commit 881a99b0a8eb4f4be924e8d6a77d4df0e12b83c2 Author: cristiano.hoshikawa Date: Fri Aug 21 08:37:51 2026 -0300 first commit diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..85e57b8 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,47 @@ +.git/ +.gitignore +.gitattributes + +__pycache__/ +*.py[cod] +*$py.class +.venv/ +venv/ +env/ +ENV/ +.pytest_cache/ +.coverage +htmlcov/ +.tox/ +.hypothesis/ + +.vscode/ +.idea/ +*.swp +*.swo +*~ +.DS_Store + +tests/ +docs/ +azure-pipelines/ +k8s/ +ddl/ +sh/ +docker-compose.yml + +.env +.env.* +!.env.example + +logs/ +timeline/ +log_agent/ +recordings/ +backup_logs_*/ +src/logs/ +src/timeline/ + +tmp/ +temp/ +*.tmp diff --git a/.env.kube.dev b/.env.kube.dev new file mode 100644 index 0000000..7fcfcf1 --- /dev/null +++ b/.env.kube.dev @@ -0,0 +1,182 @@ +# Copie para .env.dev ou .env.prod antes de executar. + +APP_ENV=dev +LOG_LEVEL=DEBUG +LOG_DIR=./logs +LOG_TO_FILE=1 +CALL_TIMELINE_DIR=./timeline +CALL_TIMELINE_ENABLED=1 +CALL_TIMELINE_CONSOLE=1 +FLOW_LOG_ENABLED=1 +FLOW_AUDIO_BURST_GAP_S=0.8 +FLOW_LOG_AUDIO_BURST_END=1 +FLOW_LOG_PREVIEW_CHARS=500 +FLOW_LOG_STT_SKIPS=1 +FLOW_LOG_VAD_DECISIONS=1 +FLOW_LOG_VAD_ACTIVITY=1 +FLOW_LOG_VAD_ACTIVITY_MIN_PROB=0.03 +STRUCTURED_EVENT_LOG_ENABLED=1 +EXPORT_DIR=./recordings +ENTIRE_CALL_RECORDING_ENABLED=1 +ENTIRE_CALL_RECORDING_TMP_DIR=./recordings/entire_call_tmp + +AGENT_BASE_NAME=ws-voice-agent +AGENT_BACKEND=remote_sse +AGENT_RECONNECT_ENABLED=1 +AGENT_RECONNECT_MAX_ATTEMPTS=1 +AGENT_RECONNECT_TIMEOUT_S=10 + +# WebSocket keepalive do Uvicorn para sessoes longas de audio. +UVICORN_WS_PING_INTERVAL_S=30 +UVICORN_WS_PING_TIMEOUT_S=120 + +LIVEKIT_URL=ws://tim-ai-atend-agnt-integ-tia-livekit.agnt-ai-atendimento-tia.svc.cluster.local:7880 +LIVEKIT_API_KEY=tia_livek_tia_api_key +LIVEKIT_API_SECRET=TiaLivekitSecret2026KeyBridgeSync01 + +# Providers locais para smoke test +STT_PROVIDER=internal_http +TTS_PROVIDER=xai +TTS_EMPTY_FRAME_RETRY_TIMEOUT_S=3 +TTS_FIRST_FRAME_TIMEOUT_S=3 +TTS_UNDERFLOW_ERROR_MS=1000 +TTS_TOTAL_TIMEOUT_S=60 +TTS_PLAYOUT_START_TIMEOUT_S=6 + +# Tuning de fala/interrupcao: permite detectar palavras curtas como "nao". +MIN_INTERRUPT_S=0.5 +DISCARD_AUDIO_IF_UNINTERRUPTIBLE=0 +SILENCE_CHECK_EVERY=1 +# Volume da fala do agente: a correcao de nivel agora e feita na saida do TTS +# (worker), com limiter soft-knee p/ nao clipar. O ganho do bridge fica neutro. +WS_OUTPUT_GAIN=1.0 +TTS_OUTPUT_GAIN=2.0 +VAD_MIN_SPEECH_DURATION=0.15 +VAD_ACTIVATION_THRESHOLD=0.30 +VAD_DEACTIVATION_THRESHOLD=0.15 +VAD_MIN_SILENCE_DURATION=1.0 +VAD_PREFIX_PADDING_DURATION=1.0 +VAD_PREFIX_PADDING_MIN_DURATION=1.0 +AGENT_WAIT_TIMEOUT_RETRY_VAD_THRESHOLD_ENABLED=1 +AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD=0.20 + +# STT fake: usa uma frase por turno. Quando acabar, repete a ultima. +FAKE_STT_TRANSCRIPTS="alô" +FAKE_STT_MODE=repeat_last +FAKE_STT_MIN_AUDIO_MS=150 + +# TTS fake: gera um tom PCM local para validar o pipeline de audio. +FAKE_TTS_TONE_HZ=440 +FAKE_TTS_AMPLITUDE=0.12 +FAKE_TTS_MIN_DURATION_MS=320 +FAKE_TTS_MAX_DURATION_MS=2200 +FAKE_TTS_CHAR_DURATION_MS=35 + +# Azure Speech TTS opcional +AZURE_TTS_IMPLEMENTATION=plugin +AZURE_SPEECH_REGION=brazilsouth +AZURE_SPEECH_ENDPOINT=https://speech-agent-ai-atendi-fqa-01.cognitiveservices.azure.com/ +AZURE_SPEECH_VOICE=pt-BR-FranciscaNeural +AZURE_SPEECH_LANGUAGE=pt-BR +AZURE_SPEECH_DEPLOYMENT_ID= + +# STT opcional +STT_URL=http://10.152.95.27:8100/api/transcriber +STT_KEY= +STT_LANG=portuguese +STT_MIN_PROB_SINGLE_WORD=0.03 +STT_MIN_AUDIO_MS=120 +STT_MIN_DBFS=-55 +VOSK_MODEL_PATH= + +# Backend remoto opcional +REMOTE_AGENT_WS_URL=ws://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/agent/ws +# Endpoint fake mantido apenas para testes manuais do contrato websocket. +REMOTE_AGENT_WS_FAKE_URL=ws://127.0.0.1:8000/fake-agent/ws +REMOTE_AGENT_WS_OPEN_TIMEOUT_S=10 +REMOTE_AGENT_WS_READ_TIMEOUT_S=900 +REMOTE_AGENT_WS_WRITE_TIMEOUT_S=10 +REMOTE_AGENT_WS_CLOSE_TIMEOUT_S=10 +REMOTE_AGENT_SSE_URL_CONTA=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/agent/sse +REMOTE_AGENT_HEALTH_URL=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/health +REMOTE_AGENT_HEALTH_URL_CONTA=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/health +REMOTE_AGENT_SSE_DEFAULT_STAGE=PRESENTATION +REMOTE_AGENT_SSE_CONNECT_TIMEOUT_S=10 +REMOTE_AGENT_SSE_READ_TIMEOUT_S=900 +REMOTE_AGENT_SSE_WRITE_TIMEOUT_S=10 +REMOTE_AGENT_INFLIGHT_WAIT_INTERVAL_S=12 +REMOTE_AGENT_INFLIGHT_WAIT_TIMEOUT_S=180 +REMOTE_AGENT_INFLIGHT_WAIT_MAX_NOTICES=0 +REMOTE_AGENT_INFLIGHT_WAIT_TEXT=Um momento, ainda estou consultando para te ajudar. +PRE_BACKEND_WAIT_NOTICE_FAST_ON_VAD_PAUSE=0 +REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR=src/app/livekit/assets/comfort/short +REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR=src/app/livekit/assets/comfort/long + +# Mock temporario para encerrar a chamada apos a primeira resposta de audio do agente. +MOCK_STOP_AFTER_FIRST_AUDIO_ENABLED=0 +MOCK_STOP_AFTER_FIRST_AUDIO_SILENCE_S=2 +MOCK_STOP_AFTER_FIRST_AUDIO_REASON=nao_resolvido + +# Bridge audio tuning +HOLD_SILENCE_DBFS=-45 +LK_CATCHUP_KEEP_MS=600 +AUDIO_IN_BACKLOG_SHED_ENABLED=1 +AUDIO_IN_BACKLOG_SHED_THRESHOLD_MS=500 +AUDIO_IN_BACKLOG_SHED_KEEP_MS=300 +AUDIO_IN_LATENCY_METRICS_ENABLED=1 +AUDIO_IN_LATENCY_ALERT_MS=1000 +AUDIO_IN_LATENCY_LOG_INTERVAL_S=15 +LIVEKIT_AUDIO_SOURCE_QUEUE_SIZE_MS=500 +LIVEKIT_AUDIO_SOURCE_CLEAR_ON_SHED=1 + +# Status terminais de finalizacao normal, resolvidos por reason no bridge. +FINAL_STOP_STATUS_RESOLVED=stop_resolvido_e_finalizado +FINAL_STOP_STATUS_UNRESOLVED=stop_nao_resolvido +FINAL_STOP_STATUS_OTHER_SUBJECT=stop_outro_assunto +FINAL_STOP_STATUS_LONG_SILENCE=stop_silencio_longo +FINAL_STOP_DEFAULT_KIND=resolved +FINAL_STOP_REASON_RESOLVED=stage_done +FINAL_STOP_REASON_UNRESOLVED=nao_resolvido +FINAL_STOP_REASON_OTHER_SUBJECT=outro_assunto +FINAL_STOP_REASON_LONG_SILENCE=no_user_response + +# Readiness / capacidade do TIA +TIA_WS_MAX_CONNECTIONS=0 +TIA_RESOURCE_HEALTH_TTL_S=5 +TIA_RESOURCE_HEALTH_TIMEOUT_S=3 +TIA_SKIP_STT_READINESS=1 +STT_HEALTH_URL=http://10.152.95.27:8100/health + +# Idle nudge +IDLE_NUDGE_ENABLED=0 +IDLE_NUDGE_DELAY_S=60 +IDLE_NUDGE_JOIN_DELAY_S=60 + +# Interrupcao durante processamento +DEFERRED_INTERRUPTION_ENABLED=1 +DEFERRED_INTERRUPTION_MIN_AUDIO_MS=1000 +DEFERRED_INTERRUPTION_STT_SETTLE_TIMEOUT_S=3.0 +DEFERRED_INTERRUPTION_USER_TURN_TIMEOUT_S=10.0 + +# XAI +XAI_TTS_VOICE=c8x2ieiocufs +XAI_TTS_LANGUAGE=pt-BR +XAI_TTS_READINESS_MODE=connect +XAI_WEBSOCKET_URL=wss://peordagnt002prd.pe.inference.generativeai.us-chicago-1.oci.oraclecloud.com/xai/v1/tts + +# PUBSUB Metrics +GCP_PROJECT_ID=tim-bigdata-dev-ca1f +AGENT_PUBSUB_TOPIC=pbs-ingest-agnt-ai-tia-events + +REMOTE_AGENT_SSE_URL_OFERTA=https://agt-ai-atendimento-ofertas-dev.internal.timbrasil.com.br/agent/execute +REMOTE_AGENT_SSE_TLS_VERIFY_OFERTA=0 + +# OTEL +OTEL_EXPORTER_OTLP_TRACES_ENDPOINT=http://10.153.35.23/v1/traces +OTEL_SERVICE_NAME=ai-agent-tia-orch +OTEL_EXPORTER_OTLP_HEADERS=Host=tim-ai-atend-agnt-opentelemetry + +#BUCKET OBJECT STORAGE +BUCKET_REGION=sa-saopaulo-1 +BUCKET_NAME=osb-gru-agnt-ai-atendimento-dev-004 +BUCKET_NAMESPACE=grfrp6rtsznm diff --git a/.env.kube.fqa b/.env.kube.fqa new file mode 100644 index 0000000..6b28fdc --- /dev/null +++ b/.env.kube.fqa @@ -0,0 +1,179 @@ +# Copie para .env.dev ou .env.prod antes de executar. + +APP_ENV=fqa +LOG_LEVEL=DEBUG +LOG_DIR=./logs +LOG_TO_FILE=1 +CALL_TIMELINE_DIR=./timeline +CALL_TIMELINE_ENABLED=1 +CALL_TIMELINE_CONSOLE=1 +FLOW_LOG_ENABLED=1 +FLOW_AUDIO_BURST_GAP_S=0.8 +FLOW_LOG_AUDIO_BURST_END=1 +FLOW_LOG_PREVIEW_CHARS=500 +FLOW_LOG_STT_SKIPS=1 +FLOW_LOG_VAD_DECISIONS=1 +FLOW_LOG_VAD_ACTIVITY=1 +FLOW_LOG_VAD_ACTIVITY_MIN_PROB=0.03 +STRUCTURED_EVENT_LOG_ENABLED=1 +EXPORT_DIR=./recordings +ENTIRE_CALL_RECORDING_ENABLED=1 +ENTIRE_CALL_RECORDING_TMP_DIR=./recordings/entire_call_tmp + +AGENT_BASE_NAME=ws-voice-agent +AGENT_BACKEND=remote_sse +AGENT_RECONNECT_ENABLED=1 +AGENT_RECONNECT_MAX_ATTEMPTS=1 +AGENT_RECONNECT_TIMEOUT_S=10 + +# WebSocket keepalive do Uvicorn para sessoes longas de audio. +UVICORN_WS_PING_INTERVAL_S=30 +UVICORN_WS_PING_TIMEOUT_S=120 + +LIVEKIT_URL=ws://tim-ai-atend-agnt-integ-tia-livekit.agnt-ai-atendimento-tia.svc.cluster.local:7880 +LIVEKIT_API_KEY=tia_livek_tia_api_key +LIVEKIT_API_SECRET=TiaLivekitSecret2026KeyBridgeSync01 + +# Providers locais para smoke test +STT_PROVIDER=internal_http +TTS_PROVIDER=xai +TTS_EMPTY_FRAME_RETRY_TIMEOUT_S=3 +TTS_FIRST_FRAME_TIMEOUT_S=3 +TTS_UNDERFLOW_ERROR_MS=1000 +TTS_TOTAL_TIMEOUT_S=60 +TTS_PLAYOUT_START_TIMEOUT_S=6 + +# Tuning de fala/interrupcao: permite detectar palavras curtas como "nao". +MIN_INTERRUPT_S=0.5 +DISCARD_AUDIO_IF_UNINTERRUPTIBLE=0 +SILENCE_CHECK_EVERY=1 +# Volume da fala do agente: a correcao de nivel agora e feita na saida do TTS +# (worker), com limiter soft-knee p/ nao clipar. O ganho do bridge fica neutro. +WS_OUTPUT_GAIN=1.0 +TTS_OUTPUT_GAIN=2.0 +VAD_MIN_SPEECH_DURATION=0.15 +VAD_ACTIVATION_THRESHOLD=0.30 +VAD_MIN_SILENCE_DURATION=1.0 +VAD_DEACTIVATION_THRESHOLD=0.15 +VAD_PREFIX_PADDING_DURATION=1.0 +VAD_PREFIX_PADDING_MIN_DURATION=1.0 +AGENT_WAIT_TIMEOUT_RETRY_VAD_THRESHOLD_ENABLED=1 +AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD=0.20 + +# STT fake: usa uma frase por turno. Quando acabar, repete a ultima. +FAKE_STT_TRANSCRIPTS="alô" +FAKE_STT_MODE=repeat_last +FAKE_STT_MIN_AUDIO_MS=150 + +# TTS fake: gera um tom PCM local para validar o pipeline de audio. +FAKE_TTS_TONE_HZ=440 +FAKE_TTS_AMPLITUDE=0.12 +FAKE_TTS_MIN_DURATION_MS=320 +FAKE_TTS_MAX_DURATION_MS=2200 +FAKE_TTS_CHAR_DURATION_MS=35 + + +# Azure Speech TTS opcional +AZURE_TTS_IMPLEMENTATION=plugin +AZURE_SPEECH_REGION=brazilsouth +AZURE_SPEECH_ENDPOINT=https://speech-agent-ai-atendi-fqa-01.cognitiveservices.azure.com/ +AZURE_SPEECH_VOICE=pt-BR-FranciscaNeural +AZURE_SPEECH_LANGUAGE=pt-BR +AZURE_SPEECH_DEPLOYMENT_ID= + +# STT opcional +STT_URL=http://10.152.95.27:8100/api/transcriber +STT_KEY= +STT_LANG=portuguese +STT_MIN_PROB_SINGLE_WORD=0.03 +STT_MIN_AUDIO_MS=120 +STT_MIN_DBFS=-55 +VOSK_MODEL_PATH= + +# Backend remoto opcional +REMOTE_AGENT_WS_URL=ws://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/agent/ws +# Endpoint fake mantido apenas para testes manuais do contrato websocket. +REMOTE_AGENT_WS_FAKE_URL=ws://127.0.0.1:8000/fake-agent/ws +REMOTE_AGENT_WS_OPEN_TIMEOUT_S=10 +REMOTE_AGENT_WS_READ_TIMEOUT_S=900 +REMOTE_AGENT_WS_WRITE_TIMEOUT_S=10 +REMOTE_AGENT_WS_CLOSE_TIMEOUT_S=10 +REMOTE_AGENT_SSE_URL_CONTA=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/agent/sse +REMOTE_AGENT_HEALTH_URL=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/health +REMOTE_AGENT_HEALTH_URL_CONTA=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/health +REMOTE_AGENT_SSE_DEFAULT_STAGE=PRESENTATION +REMOTE_AGENT_SSE_CONNECT_TIMEOUT_S=10 +REMOTE_AGENT_SSE_READ_TIMEOUT_S=900 +REMOTE_AGENT_SSE_WRITE_TIMEOUT_S=10 +REMOTE_AGENT_INFLIGHT_WAIT_INTERVAL_S=12 +REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR=src/app/livekit/assets/comfort/short +REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR=src/app/livekit/assets/comfort/long +REMOTE_AGENT_INFLIGHT_WAIT_TIMEOUT_S=180 +REMOTE_AGENT_INFLIGHT_WAIT_MAX_NOTICES=0 +REMOTE_AGENT_INFLIGHT_WAIT_TEXT=Um momento, ainda estou consultando para te ajudar. +PRE_BACKEND_WAIT_NOTICE_FAST_ON_VAD_PAUSE=0 + +# Mock temporario para encerrar a chamada apos a primeira resposta de audio do agente. +MOCK_STOP_AFTER_FIRST_AUDIO_ENABLED=0 +MOCK_STOP_AFTER_FIRST_AUDIO_SILENCE_S=2 +MOCK_STOP_AFTER_FIRST_AUDIO_REASON=nao_resolvido + +# Bridge audio tuning +HOLD_SILENCE_DBFS=-45 +LK_CATCHUP_KEEP_MS=600 +AUDIO_IN_BACKLOG_SHED_ENABLED=1 +AUDIO_IN_BACKLOG_SHED_THRESHOLD_MS=500 +AUDIO_IN_BACKLOG_SHED_KEEP_MS=300 +AUDIO_IN_LATENCY_METRICS_ENABLED=1 +AUDIO_IN_LATENCY_ALERT_MS=1000 +AUDIO_IN_LATENCY_LOG_INTERVAL_S=15 +LIVEKIT_AUDIO_SOURCE_QUEUE_SIZE_MS=500 +LIVEKIT_AUDIO_SOURCE_CLEAR_ON_SHED=1 + +# Status terminais de finalizacao normal, resolvidos por reason no bridge. +FINAL_STOP_STATUS_RESOLVED=stop_resolvido_e_finalizado +FINAL_STOP_STATUS_UNRESOLVED=stop_nao_resolvido +FINAL_STOP_STATUS_OTHER_SUBJECT=stop_outro_assunto +FINAL_STOP_STATUS_LONG_SILENCE=stop_silencio_longo +FINAL_STOP_DEFAULT_KIND=resolved +FINAL_STOP_REASON_RESOLVED=stage_done +FINAL_STOP_REASON_UNRESOLVED=nao_resolvido +FINAL_STOP_REASON_OTHER_SUBJECT=outro_assunto +FINAL_STOP_REASON_LONG_SILENCE=no_user_response + +# Readiness / capacidade do TIA +TIA_WS_MAX_CONNECTIONS=0 +TIA_RESOURCE_HEALTH_TTL_S=5 +TIA_RESOURCE_HEALTH_TIMEOUT_S=3 +TIA_SKIP_STT_READINESS=1 +STT_HEALTH_URL=http://10.152.95.27:8100/health + +# Idle nudge +IDLE_NUDGE_ENABLED=0 +IDLE_NUDGE_DELAY_S=60 +IDLE_NUDGE_JOIN_DELAY_S=60 + +# Interrupcao durante processamento +DEFERRED_INTERRUPTION_ENABLED=1 +DEFERRED_INTERRUPTION_MIN_AUDIO_MS=1000 +DEFERRED_INTERRUPTION_STT_SETTLE_TIMEOUT_S=3.0 +DEFERRED_INTERRUPTION_USER_TURN_TIMEOUT_S=10.0 + +# XAI +XAI_TTS_VOICE=c8x2ieiocufs +XAI_TTS_LANGUAGE=pt-BR +XAI_TTS_READINESS_MODE=connect +XAI_WEBSOCKET_URL=wss://peordagnt002prd.pe.inference.generativeai.us-chicago-1.oci.oraclecloud.com/xai/v1/tts + +GCP_PROJECT_ID=tim-bigdata-fqa-60ce +AGENT_PUBSUB_TOPIC=pbs-ingest-agnt-ai-tia-events + +# OTEL +OTEL_EXPORTER_OTLP_TRACES_ENDPOINT=http://10.153.35.71/v1/traces +OTEL_SERVICE_NAME=ai-agent-tia-orch +OTEL_EXPORTER_OTLP_HEADERS=Host=tim-ai-atend-agnt-opentelemetry + +#BUCKET OBJECT STORAGE +BUCKET_REGION=sa-saopaulo-1 +BUCKET_NAME=osb-gru-agnt-ai-atendimento-fqa-004 +BUCKET_NAMESPACE=grfrp6rtsznm diff --git a/.env.kube.prd b/.env.kube.prd new file mode 100644 index 0000000..e4edd4d --- /dev/null +++ b/.env.kube.prd @@ -0,0 +1,175 @@ +# Copie para .env.dev ou .env.prod antes de executar. + +APP_ENV=prd +LOG_LEVEL=INFO +LOG_DIR=./logs +CALL_TIMELINE_DIR=./timeline +CALL_TIMELINE_CONSOLE=0 +FLOW_LOG_ENABLED=1 +FLOW_AUDIO_BURST_GAP_S=0.8 +FLOW_LOG_AUDIO_BURST_END=1 +FLOW_LOG_PREVIEW_CHARS=500 +FLOW_LOG_STT_SKIPS=1 +FLOW_LOG_VAD_DECISIONS=1 +FLOW_LOG_VAD_ACTIVITY=1 +FLOW_LOG_VAD_ACTIVITY_MIN_PROB=0.03 +EXPORT_DIR=./recordings +ENTIRE_CALL_RECORDING_ENABLED=1 +ENTIRE_CALL_RECORDING_TMP_DIR=./recordings/entire_call_tmp + +AGENT_BASE_NAME=ws-voice-agent +AGENT_BACKEND=remote_sse +AGENT_RECONNECT_ENABLED=1 +AGENT_RECONNECT_MAX_ATTEMPTS=1 +AGENT_RECONNECT_TIMEOUT_S=10 + +# WebSocket keepalive do Uvicorn para sessoes longas de audio. +UVICORN_WS_PING_INTERVAL_S=30 +UVICORN_WS_PING_TIMEOUT_S=120 + +LIVEKIT_URL=ws://tim-ai-atend-agnt-integ-tia-livekit.agnt-ai-atendimento-tia.svc.cluster.local:7880 +LIVEKIT_API_KEY=tia_livek_tia_api_key +LIVEKIT_API_SECRET=TiaLivekitSecret2026KeyBridgeSync01 + +# Providers locais para smoke test +STT_PROVIDER=internal_http +TTS_PROVIDER=xai +TTS_EMPTY_FRAME_RETRY_TIMEOUT_S=3 +TTS_FIRST_FRAME_TIMEOUT_S=3 +TTS_UNDERFLOW_ERROR_MS=1000 +TTS_TOTAL_TIMEOUT_S=60 + +# Tuning de fala/interrupcao: permite detectar palavras curtas como "nao". +MIN_INTERRUPT_S=0.5 +DISCARD_AUDIO_IF_UNINTERRUPTIBLE=0 +SILENCE_CHECK_EVERY=1 +# Volume da fala do agente: a correcao de nivel agora e feita na saida do TTS +# (worker), com limiter soft-knee p/ nao clipar. O ganho do bridge fica neutro. +WS_OUTPUT_GAIN=1.0 +TTS_OUTPUT_GAIN=2.0 +VAD_MIN_SPEECH_DURATION=0.15 +VAD_ACTIVATION_THRESHOLD=0.30 +VAD_MIN_SILENCE_DURATION=1.0 +VAD_DEACTIVATION_THRESHOLD=0.15 +VAD_PREFIX_PADDING_DURATION=1.0 +VAD_PREFIX_PADDING_MIN_DURATION=1.0 +AGENT_WAIT_TIMEOUT_RETRY_VAD_THRESHOLD_ENABLED=1 +AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD=0.20 + +# STT fake: usa uma frase por turno. Quando acabar, repete a ultima. +FAKE_STT_TRANSCRIPTS="alô" +FAKE_STT_MODE=repeat_last +FAKE_STT_MIN_AUDIO_MS=150 + +# TTS fake: gera um tom PCM local para validar o pipeline de audio. +FAKE_TTS_TONE_HZ=440 +FAKE_TTS_AMPLITUDE=0.12 +FAKE_TTS_MIN_DURATION_MS=320 +FAKE_TTS_MAX_DURATION_MS=2200 +FAKE_TTS_CHAR_DURATION_MS=35 + + +# Azure Speech TTS opcional +AZURE_TTS_IMPLEMENTATION=plugin +AZURE_SPEECH_REGION=brazilsouth +AZURE_SPEECH_ENDPOINT= +AZURE_SPEECH_VOICE=pt-BR-FranciscaNeural +AZURE_SPEECH_LANGUAGE=pt-BR +AZURE_SPEECH_DEPLOYMENT_ID= + +# STT opcional +STT_URL=http://10.152.89.211:8100/api/transcriber +STT_KEY= +STT_LANG=portuguese +STT_MIN_PROB_SINGLE_WORD=0.03 +STT_MIN_AUDIO_MS=120 +STT_MIN_DBFS=-55 +VOSK_MODEL_PATH= + +# Backend remoto opcional +REMOTE_AGENT_WS_URL=ws://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/agent/ws +# Endpoint fake mantido apenas para testes manuais do contrato websocket. +REMOTE_AGENT_WS_FAKE_URL=ws://127.0.0.1:8000/fake-agent/ws +REMOTE_AGENT_WS_OPEN_TIMEOUT_S=10 +REMOTE_AGENT_WS_READ_TIMEOUT_S=900 +REMOTE_AGENT_WS_WRITE_TIMEOUT_S=10 +REMOTE_AGENT_WS_CLOSE_TIMEOUT_S=10 +REMOTE_AGENT_SSE_URL_CONTA=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/agent/sse +REMOTE_AGENT_HEALTH_URL=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/health +REMOTE_AGENT_HEALTH_URL_CONTA=http://tim-ai-contas-agnt-service.agnt-ai-atendimento-contas.svc.cluster.local:80/health +REMOTE_AGENT_SSE_DEFAULT_STAGE=PRESENTATION +REMOTE_AGENT_SSE_CONNECT_TIMEOUT_S=10 +REMOTE_AGENT_SSE_READ_TIMEOUT_S=900 +REMOTE_AGENT_SSE_WRITE_TIMEOUT_S=10 +REMOTE_AGENT_INFLIGHT_WAIT_INTERVAL_S=12 +REMOTE_AGENT_INFLIGHT_WAIT_TIMEOUT_S=180 +REMOTE_AGENT_INFLIGHT_WAIT_MAX_NOTICES=0 +REMOTE_AGENT_INFLIGHT_WAIT_TEXT=Um momento, ainda estou consultando para te ajudar. +PRE_BACKEND_WAIT_NOTICE_FAST_ON_VAD_PAUSE=0 +REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR=src/app/livekit/assets/comfort/short +REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR=src/app/livekit/assets/comfort/long + +# Mock temporario para encerrar a chamada apos a primeira resposta de audio do agente. +MOCK_STOP_AFTER_FIRST_AUDIO_ENABLED=0 +MOCK_STOP_AFTER_FIRST_AUDIO_SILENCE_S=2 +MOCK_STOP_AFTER_FIRST_AUDIO_REASON=nao_resolvido + +# Bridge audio tuning +HOLD_SILENCE_DBFS=-45 +LK_CATCHUP_KEEP_MS=600 +AUDIO_IN_BACKLOG_SHED_ENABLED=1 +AUDIO_IN_BACKLOG_SHED_THRESHOLD_MS=500 +AUDIO_IN_BACKLOG_SHED_KEEP_MS=300 +AUDIO_IN_LATENCY_METRICS_ENABLED=1 +AUDIO_IN_LATENCY_ALERT_MS=1000 +AUDIO_IN_LATENCY_LOG_INTERVAL_S=15 +LIVEKIT_AUDIO_SOURCE_QUEUE_SIZE_MS=500 +LIVEKIT_AUDIO_SOURCE_CLEAR_ON_SHED=1 + +# Status terminais de finalizacao normal, resolvidos por reason no bridge. +FINAL_STOP_STATUS_RESOLVED=stop_resolvido_e_finalizado +FINAL_STOP_STATUS_UNRESOLVED=stop_nao_resolvido +FINAL_STOP_STATUS_OTHER_SUBJECT=stop_outro_assunto +FINAL_STOP_STATUS_LONG_SILENCE=stop_silencio_longo +FINAL_STOP_DEFAULT_KIND=resolved +FINAL_STOP_REASON_RESOLVED=stage_done +FINAL_STOP_REASON_UNRESOLVED=nao_resolvido +FINAL_STOP_REASON_OTHER_SUBJECT=outro_assunto +FINAL_STOP_REASON_LONG_SILENCE=no_user_response + +# Readiness / capacidade do TIA +TIA_WS_MAX_CONNECTIONS=0 +TIA_RESOURCE_HEALTH_TTL_S=5 +TIA_RESOURCE_HEALTH_TIMEOUT_S=3 +TIA_SKIP_STT_READINESS=1 +STT_HEALTH_URL=http://10.152.89.211:8100/healthz + +# Idle nudge +IDLE_NUDGE_ENABLED=0 +IDLE_NUDGE_DELAY_S=60 +IDLE_NUDGE_JOIN_DELAY_S=60 + +# Interrupcao durante processamento +DEFERRED_INTERRUPTION_ENABLED=1 +DEFERRED_INTERRUPTION_MIN_AUDIO_MS=1000 +DEFERRED_INTERRUPTION_STT_SETTLE_TIMEOUT_S=3.0 +DEFERRED_INTERRUPTION_USER_TURN_TIMEOUT_S=10.0 + +# XAI +XAI_TTS_VOICE=c8x2ieiocufs +XAI_TTS_LANGUAGE=pt-BR +XAI_TTS_READINESS_MODE=connect +XAI_WEBSOCKET_URL=wss://peordagnt002prd.pe.inference.generativeai.us-chicago-1.oci.oraclecloud.com/xai/v1/tts + +GCP_PROJECT_ID=tim-bigdata-prod-e305 +AGENT_PUBSUB_TOPIC=pbs-ingest-agnt-ai-tia-events + +# OTEL +OTEL_EXPORTER_OTLP_TRACES_ENDPOINT=http://10.152.6.254/v1/traces +OTEL_SERVICE_NAME=ai-agent-tia-orch +OTEL_EXPORTER_OTLP_HEADERS=Host=tim-ai-atend-agnt-opentelemetry + +#BUCKET OBJECT STORAGE +BUCKET_REGION=sa-saopaulo-1 +BUCKET_NAME=osb-gru-agnt-ai-atendimento-prd-004 +BUCKET_NAMESPACE=grfrp6rtsznm diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..68cfe57 --- /dev/null +++ b/.gitignore @@ -0,0 +1,54 @@ +# OS / editor +.DS_Store +.idea/ +.vscode/ +*.iml +*.swp +*.swo +*~ + +# Python +__pycache__/ +*.py[cod] +*$py.class +.python-version +.mypy_cache/ +.pytest_cache/ +.hypothesis/ +.tox/ +.coverage +.coverage.* +htmlcov/ +build/ +dist/ +*.egg-info/ + +# Virtual environments +.venv/ +venv/ +env/ +ENV/ + +# Environment files +.env +.env.dev +!.env.example + +# Runtime artifacts +logs/ +timeline/ +log_agent/ +recordings/ +backup_logs_*/ +src/logs/ +src/timeline/ +.run/ + +# Documentation / temporary +docs/_build/ +tmp/ +temp/ +*.tmp + +# Cache +cache/ diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..fd1301d --- /dev/null +++ b/Dockerfile @@ -0,0 +1,26 @@ +FROM python:3.12-slim + +WORKDIR /app + +ENV PYTHONDONTWRITEBYTECODE=1 +ENV PYTHONUNBUFFERED=1 +ENV PYTHONPATH=/app/src + +RUN apt-get update && apt-get install -y --no-install-recommends \ + build-essential \ + curl \ + libsndfile1 \ +&& rm -rf /var/lib/apt/lists/* + +COPY requirements.txt ./ +RUN pip install uv && \ + uv pip install --no-cache-dir --system -r requirements.txt + +COPY src/ ./src/ + +RUN useradd -m -u 1000 agent && chown -R agent:agent /app +USER agent + +EXPOSE 8000 18081 + +ENTRYPOINT ["python", "-m"] diff --git a/README.md b/README.md new file mode 100644 index 0000000..4e1c84a --- /dev/null +++ b/README.md @@ -0,0 +1,253 @@ +# TIA Voice API + +Este repositorio contem a stack de voz-to-voz do projeto TIA. + +## Estrutura + +- `src/app`: gateway WebSocket, runtime LiveKit, providers e utilitarios +- `src/agent`: pipeline e estagios do agente +- `tests`: testes automatizados +- `docs`: documentacao viva do projeto +- `k8s/livekit`: imagem e manifest do pod dedicado do LiveKit +- `k8s/tia`: imagem e manifest do pod da aplicacao com `bridge` e `agent` +- `requirements.txt`: dependencias Python da aplicacao +- `livekit.yaml`: configuracao local do servidor LiveKit + +## Setup local + +```bash +make setup +``` + +O alvo `make setup` faz o bootstrap do ambiente local: + +- cria `.venv` priorizando `python3.13`, `python3.12`, `python3.11`, `python3.10`, `python3.9`, `python3` e `python` +- instala as dependencias de `requirements.txt` +- cria `.env.dev` a partir de `.env.example` se o arquivo ainda nao existir + +O projeto hoje exige Python `3.9` ate `3.13`. Python `3.14` nao e aceito pelas dependencias atuais de `livekit-agents==1.3.10`. + +Se voce tiver mais de um Python instalado e quiser forcar uma versao especifica no bootstrap: + +```bash +make BOOTSTRAP_PY=python3.12 setup +``` + +Se voce preferir fazer manualmente ou quiser depurar o bootstrap, use: + +```bash +python3.12 -m venv .venv +source .venv/bin/activate +python -m pip install --upgrade pip +python -m pip install -r requirements.txt +cp .env.example .env.dev +``` + +Se `make test`, `make agent` ou `make bridge` falharem com erro de virtualenv ausente, rode `make setup` primeiro. + +## Teste local + +```bash +make test +``` + +O `makefile` executa os testes com `.venv/bin/python` e injeta `PYTHONPATH=src`, entao o comando deve ser chamado a partir da raiz do repositorio. + +## Execucao local + +O `.env.dev` de desenvolvimento pode ser configurado em dois modos: + +- smoke test local: `AGENT_BACKEND=remote_ws_fake`, `STT_PROVIDER=fake` e `TTS_PROVIDER=fake` +- integracao real: ajuste `AGENT_BACKEND`, `STT_PROVIDER`, `TTS_PROVIDER` e as credenciais externas necessarias + +No modo fake, o agent nao depende de STT HTTP, ElevenLabs ou backend LLM externo. O `STT fake` consome uma sequencia configurada em `FAKE_STT_TRANSCRIPTS` e o `TTS fake` gera um tom PCM local para validar o pipeline de audio. + +Importante: `LIVEKIT_API_KEY` e `LIVEKIT_API_SECRET` do `.env.dev` precisam ser identicos ao bloco `keys` de `livekit.yaml`. Se esses valores divergirem, o bridge falha ao conectar/publicar no LiveKit e o dispatch do agent retorna erro de autenticacao. + +Para usar o STT HTTP real, configure no `.env.dev`: + +- `STT_PROVIDER=internal_http` +- `STT_URL=http://10.152.95.27:8100` +- `STT_KEY=...` se o servico exigir o header `x-api-key` +- `STT_LANG=portuguese` + +Observacao: o provider faz `POST` exatamente na URL configurada em `STT_URL`. Se o seu servico expuser uma rota especifica, use a URL completa, por exemplo `http://10.152.95.27:8100/transcribe`. + +Para usar Azure Speech TTS com o plugin oficial do LiveKit, configure no `.env.dev`: + +- `TTS_PROVIDER=azure` +- `AZURE_SPEECH_KEY=...` +- `AZURE_SPEECH_REGION=...` ou `AZURE_SPEECH_ENDPOINT=https://...cognitiveservices.azure.com/` +- `AZURE_SPEECH_VOICE=pt-BR-FranciscaNeural` +- `AZURE_SPEECH_LANGUAGE=pt-BR` +- `AZURE_SPEECH_DEPLOYMENT_ID=...` apenas para Custom Voice + +Para usar xAI TTS, configure no `.env.dev`: + +- `TTS_PROVIDER=xai` +- `XAI_API_KEY=...` +- `XAI_WEBSOCKET_URL=https://cloud9.api.x.ai/v1/tts` +- `XAI_TTS_READINESS_MODE=connect` opcional para validar apenas o handshake do WebSocket; default `synthesize` +- `XAI_TTS_VOICE=ara` opcional +- `XAI_TTS_LANGUAGE=pt-BR` opcional +- `TTS_FRAME_GAP_TIMEOUT_S=2` limite entre deltas de audio do xAI; quando estoura, o runtime registra `Falha TTS`, toca o audio de conforto e reenvia o texto + +Para enviar logs estruturados ao Google Pub/Sub, configure: + +- `GCP_PROJECT_ID=...` +- `AGENT_PUBSUB_TOPIC=agent-logs` ou `projects/.../topics/agent-logs` + +Se uma dessas variaveis nao existir, os eventos estruturados continuam indo apenas para o log local. O ambiente tambem precisa ter credenciais Google disponiveis via Application Default Credentials ou service account com permissao de publicacao no topico. + +Para subir o fluxo local de voz: + +```bash +make local-up +``` + +Depois abra `http://127.0.0.1:8000/voice-client`. O cliente web agora vem com `remote_ws_fake`, `fake` STT e `fake` TTS por padrao para smoke test local. + +Para rodar o teste operacional com STT Sofya, TTS xAI, Bridge e LiveKit reais: + +```bash +make local-stresstest +``` + +O alvo sobe o ambiente local se necessario, executa `python -m app.tools.local_stresstest` e gera os artefatos em `.run/local-stresstest/`: + +- `report.md`: resumo em Markdown com tabelas e diagramas Mermaid +- `summary.json`: resultado estruturado completo +- `stt_results.csv`, `tts_results.csv`, `e2e_results.csv`: planilhas dos cenarios +- `timeline_excerpt.jsonl`: eventos relevantes da timeline da chamada E2E +- `diagrams/*.svg`: imagens estaticas dos fluxos, boas para preview no VSCode +- `diagrams/*.mmd`: fontes Mermaid dos fluxos +- `stt_dumps/*.wav` e `stt_dumps/*.pcm`: audio efetivamente enviado ao STT +- WAVs gerados para baseline, variacoes, TTS e audio recebido do Bridge + +Por padrao, o alvo habilita dumps de audio do STT e logs de VAD do agent +(`FLOW_LOG_VAD_DECISIONS=1`, `FLOW_LOG_VAD_ACTIVITY=1`). O agent local e reiniciado +antes do teste para garantir que esses flags entrem no processo. + +Se `STRESS_AUDIO` nao for informado, o runner gera um baseline sintetico com xAI TTS usando `STRESS_SYNTH_TEXT`, ou `STRESS_EXPECTED_TEXT` quando `STRESS_SYNTH_TEXT` nao for definido. Quando houver uma gravacao humana, rode: + +```bash +STRESS_AUDIO=/caminho/audio_16k_mono.wav make local-stresstest +``` + +Variaveis uteis: + +- `STRESS_EXPECTED_TEXT`: frase esperada para calculo de WER +- `STRESS_SYNTH_TEXT`: frase usada na sintese do baseline xAI TTS, default igual a `STRESS_EXPECTED_TEXT` +- `STRESS_CRITICAL_TERMS`: termos obrigatorios separados por virgula +- `STRESS_WER_THRESHOLD`: limite de WER, default `0.20` +- `STRESS_REPEAT`: repeticoes dos cenarios STT, default `1`; com os 25 cenarios atuais, gera 25 chamadas ao STT. Use `STRESS_REPEAT=2` para 50 chamadas +- `STRESS_CONCURRENCY`: concorrencia dos cenarios STT, default `2` +- `STRESS_PREFIX_TEXT`: prefixo que precisa aparecer no inicio da transcricao, default primeiras palavras de `STRESS_EXPECTED_TEXT` +- `STRESS_PREFIX_WORDS`: quantidade de palavras usadas no prefixo automatico, default `2` +- `STRESS_VAD_PROXY_DBFS`: limiar RMS usado na analise local de VAD proxy, default `-45` +- `STRESS_VAD_PREFIX_PADDING_MS`: padding usado na analise local de VAD proxy, default vem de `VAD_PREFIX_PADDING_DURATION` +- `STRESS_VAD_MIN_SPEECH_MS`: fala minima usada na analise local de VAD proxy, default `100` +- `VAD_PREFIX_PADDING_DURATION`: pre-roll do LiveKit VAD, default `2.0` segundos +- `VAD_PREFIX_PADDING_MIN_DURATION`: piso aplicado tambem sobre override de `call_config`, default `2.0` segundos +- `STT_INPUT_PREFIX_PADDING_MS`: silencio curto adicionado antes do WAV enviado ao STT HTTP, default `250` +- `STT_DUMP_DIR`: diretorio dos WAV/PCM enviados ao STT, default `.run/local-stresstest/stt_dumps` +- `FLOW_LOG_VAD_DECISIONS`: habilita logs `vad_speech_start`/`vad_speech_end`, default `1` no alvo +- `FLOW_LOG_VAD_ACTIVITY`: habilita logs `vad_activity`, default `1` no alvo +- `STRESS_RESTART_AGENT_FOR_DIAGNOSTICS`: reinicia o agent antes do teste, default `1` +- `STRESS_STARTUP_WAIT_S`: tempo maximo para aguardar Bridge e agent runtime, default `180` +- `STRESS_STARTUP_POLL_S`: intervalo entre probes de prontidao, default `3` +- `STRESS_SKIP_LOCAL_WAIT`: use `1` para pular a espera inicial +- `STRESS_BRIDGE_HEALTH_URL`: URL de health do Bridge, derivada de `STRESS_BRIDGE_URL` por padrao +- `STRESS_AGENT_HEALTH_URL`: URL de health do agent runtime, default `http://127.0.0.1:18081/` +- `STRESS_REPORT_DIR`: diretorio de saida, default `.run/local-stresstest` + +Comandos auxiliares do modo local: + +```bash +make local-status +make local-logs +make bridge-logs +make agent-logs +make local-down +``` + +Se o `agent` falhar com `worker process is not responding.. worker crashed?` e o log mostrar erro de bind na porta `18081`, existe um worker antigo preso nessa porta. Nesse caso, rode `make agent-down` e depois `make agent-up` ou `make local-up` novamente. + +## Kubernetes + +A pasta `k8s` agora esta organizada em dois blocos: + +- `k8s/livekit/Dockerfile`: usa a imagem oficial `livekit/livekit-server` como base para publicacao no registro privado da empresa +- `k8s/livekit/deployment.yaml`: deployment e service do pod exclusivo do LiveKit +- `k8s/tia/Dockerfile`: mesma imagem da aplicacao Python da raiz do repo, mantida ao lado do manifest para facilitar pipeline de build/publicacao +- `k8s/tia/deployment.yaml`: deployment e service do pod `tia-app`, com dois containers na mesma imagem: `tia-bridge` e `tia-agent` + +Os dois `Dockerfile`s devem ser buildados a partir da raiz do repositorio, porque dependem do contexto raiz para copiar `requirements.txt`, `src/` e `livekit.yaml`. + +Exemplos: + +```bash +docker build -f k8s/livekit/Dockerfile -t registry.example.com/tia/livekit:TAG . +docker build -f k8s/tia/Dockerfile -t registry.example.com/tia/app:TAG . +``` + +Depois do push para o registry privado, ajuste as imagens em `k8s/livekit/deployment.yaml` e `k8s/tia/deployment.yaml` ou substitua os placeholders via pipeline: + +- `LIVEKIT_IMAGE_REPOSITORY:LIVEKIT_IMAGE_TAG` +- `TIA_IMAGE_REPOSITORY:TIA_IMAGE_TAG` + +Aplicacao dos manifests: + +```bash +kubectl apply -f k8s/livekit/deployment.yaml +kubectl apply -f k8s/tia/deployment.yaml +``` + +Observacoes: + +- `k8s/livekit/Dockerfile` embute o arquivo `livekit.yaml` dentro da imagem; se a configuracao do LiveKit mudar, a imagem precisa ser rebuildada +- `k8s/tia/deployment.yaml` ainda usa envs inline de placeholder para `LIVEKIT_API_KEY` e `LIVEKIT_API_SECRET`; o passo seguinte natural e migrar isso para `Secret` e demais configs para `ConfigMap` +- o `bridge` atende na porta `8000` e o `agent` expõe a porta interna `18081` + +## Comandos uteis + +```bash +make test +make agent +make bridge +make livekit +``` + +Todos os comandos assumem o codigo em `src`, portanto a raiz do projeto deve ser usada como diretorio de execucao. + +## Documentacao + +- `docs/README.md`: indice das docs do projeto +- `docs/refactor-plan.md`: plano do refactor incremental +- `docs/refactor-log.md`: registro das mudancas por etapa +- `docs/api-overview.md`: documentacao viva da API e do comportamento atual + +--- + +## Regional xAI TTS pool (Kubernetes) + +Esta versão inclui uma arquitetura opcional para alta disponibilidade do TTS xAI com Pods TIA regionais, pool WebSocket pré-aquecido por Pod, readiness orientada a capacidade e failover via Service/Load Balancer. + +Documentação: + +- `docs/regional/ARQUITETURA_TIA_XAI_REGIONAL.md` +- `docs/regional/DEPLOYMENT_TIA_XAI_REGIONAL.md` +- `docs/regional/TESTES_TIA_XAI_REGIONAL.md` + +Código principal: + +- `src/app/livekit/adapters/xai_pool_proxy.py` + +Manifests e scripts: + +- `k8s/regional/` +- `scripts/render-regional-k8s.sh` +- `scripts/validate-regional-k8s.sh` +- `scripts/deploy-regional-k8s.sh` + +O modo antigo de conexão direta com xAI foi preservado. O deployment regional é opt-in. diff --git a/azure-pipelines/dev.yml b/azure-pipelines/dev.yml new file mode 100644 index 0000000..ffdd0d2 --- /dev/null +++ b/azure-pipelines/dev.yml @@ -0,0 +1,136 @@ +trigger: + branches: + include: + - develop + # paths: + # exclude: + # - azure-pipelines/* + +pool: + name: Azure Pipelines + vmImage: 'ubuntu-latest' + +resources: + repositories: + - repository: templates + type: git + name: Data_Science_AI_ML/template-yaml-devops-mlops + ref: master + +variables: +- template: oke_agents_atendimento/variables.yaml@templates +- group: sa-timsa-it-agnt-ai-atendimento-cicd-dev-auth-token +- group: tim-ai-atend-agnt-integ-tia-dev-fqa +- name: imageRepository + value: '$(ociRegistry)/$(OCI_NAMESPACE)/agnt-ai-atendimento/dev/agnt-ai-atendimento-tia' +- name: imageLivekit + value: '$(ociRegistry)/$(OCI_NAMESPACE)/agnt-ai-atendimento/dev/agnt-ai-atendimento-tia-livekit' +- name: DNS + value: agt-ai-atendimento-tia-dev.internal.timbrasil.com.br +- name: REDIS_IP + value: 10.152.96.51 +- name: REDIS_HOST + value: aaaaehl73aayzoctg2kvzc5xphvuiryrtwelktfaxcc3ill4cs7ml5a-p.redis.sa-saopaulo-1.oci.oraclecloud.com +# Livekit Config +- name: LIVEKIT_REPLICAS + value: 2 +- name: CPU_LIVEKIT_REQ + value: 1 +- name: CPU_LIVEKIT_LIM + value: 2 +- name: MEM_LIVEKIT_REQ + value: 2Gi +- name: MEM_LIVEKIT_LIM + value: 4Gi +# TIA Config +- name: TIA_REPLICAS + value: 2 +- name: CPU_TIA_REQ + value: 1 +- name: CPU_TIA_LIM + value: 2 +- name: MEM_TIA_REQ + value: 4Gi +- name: MEM_TIA_LIM + value: 6Gi +# TIA Bridge Config +- name: CPU_TIA_BRIDGE_REQ + value: 1 +- name: CPU_TIA_BRIDGE_LIM + value: 2 +- name: MEM_TIA_BRIDGE_REQ + value: 2Gi +- name: MEM_TIA_BRIDGE_LIM + value: 4Gi + +name: 0.0.$(Rev:r) + +stages: +- stage: Build + jobs: + - job: Build + steps: + - script: | + echo "$(auth-token)" | docker login $(ociRegistry) -u "$(OCI_NAMESPACE)/$(SA_NAME)" --password-stdin + displayName: 'Login to OCI Registry' + + - script: | + set -eo pipefail + + docker build -f $(Build.SourcesDirectory)/k8s/tia/Dockerfile -t $(imageRepository):$(Build.BuildNumber) $(Build.SourcesDirectory) + + docker build -f $(Build.SourcesDirectory)/k8s/livekit/Dockerfile -t $(imageLivekit):$(Build.BuildNumber) $(Build.SourcesDirectory) + displayName: 'Build the image' + + - script: | + docker push $(imageRepository):$(Build.BuildNumber) + docker push $(imageLivekit):$(Build.BuildNumber) + displayName: 'Push image to OCI Registry' + +- stage: Deploy + dependsOn: Build + condition: succeeded() + jobs: + - job: deployb8s + pool: + name: oci-kubernetes + variables: + K8S_NAMESPACE: 'agnt-ai-atendimento-tia' + IMAGE_REPOSITORY: '$(imageRepository)' + IMAGE_REPOSITORY_LIVEKIT: '$(imageLivekit)' + IMAGE_TAG: '$(Build.BuildNumber)' + APP_NAME: '$(Build.Repository.Name)' + OCI_CLI_CONFIG_FILE: '$(Pipeline.Workspace)/oci_config' + + steps: + + - template: oke_agents_atendimento/autenticacao_kubernetes_v2.yaml@templates + + - template: oke_agents_atendimento/apply_env_as_configmap.yaml@templates + + - template: oke_agents_atendimento/sa_gcp_secret.yaml@templates + + - script: | + set -eo pipefail + envsubst < k8s/secrets.yaml | kubectl apply -n $(K8S_NAMESPACE) -f - + displayName: "aplicando secrets" + env: + AZURE_SPEECH_KEY: $(AZURE_SPEECH_KEY) + XAI_API_KEY: $(XAI_API_KEY_DEV) + + - script: | + set -eo pipefail + + # cada pipeline que vamos querer deployar + mkdir -p helm_build/templates + envsubst < k8s/livekit/configmap.yaml > helm_build/templates/configmap-livekit.yaml + envsubst < k8s/livekit/deployment.yaml > helm_build/templates/deployment-livekit.yaml + envsubst < k8s/tia/deployment.yaml > helm_build/templates/deployment-tia.yaml + + envsubst < k8s/helm_deploy.yaml > helm_build/Chart.yaml + helm upgrade --install $(APP_NAME) ./helm_build \ + --namespace $(K8S_NAMESPACE) \ + --wait \ + --timeout 270s \ + --atomic + displayName: 'Apply kube custom helm' diff --git a/azure-pipelines/fqa.yml b/azure-pipelines/fqa.yml new file mode 100644 index 0000000..f481d3d --- /dev/null +++ b/azure-pipelines/fqa.yml @@ -0,0 +1,142 @@ +trigger: + branches: + include: + - fqa + +pool: + name: Azure Pipelines + vmImage: 'ubuntu-latest' + +resources: + repositories: + - repository: templates + type: git + name: Data_Science_AI_ML/template-yaml-devops-mlops + ref: master + +variables: +- template: oke_agents_atendimento/variables.yaml@templates +- group: sa-timsa-it-agnt-ai-atendimento-cicd-fqa-auth-token +- group: tim-ai-atend-agnt-integ-tia-dev-fqa +- name: imageRepository + value: '$(ociRegistry)/$(OCI_NAMESPACE)/agnt-ai-atendimento/fqa/agnt-ai-atendimento-tia' +- name: imageLivekit + value: '$(ociRegistry)/$(OCI_NAMESPACE)/agnt-ai-atendimento/fqa/agnt-ai-atendimento-tia-livekit' +- name: DNS + value: agt-ai-atendimento-tia-fqa.internal.timbrasil.com.br +- name: REDIS_IP + value: 10.152.96.51 +- name: REDIS_HOST + value: aaaaehl73aayzoctg2kvzc5xphvuiryrtwelktfaxcc3ill4cs7ml5a-p.redis.sa-saopaulo-1.oci.oraclecloud.com + +# Livekit Config +- name: LIVEKIT_REPLICAS + value: 2 +- name: CPU_LIVEKIT_REQ + value: 1 +- name: CPU_LIVEKIT_LIM + value: 2 +- name: MEM_LIVEKIT_REQ + value: 2Gi +- name: MEM_LIVEKIT_LIM + value: 4Gi + +# TIA Config +- name: TIA_REPLICAS + value: 2 +# esse valor de réplica vale para ambos os abaixo +- name: CPU_TIA_REQ + value: 1 +- name: CPU_TIA_LIM + value: 2 +- name: MEM_TIA_REQ + value: 4Gi +- name: MEM_TIA_LIM + value: 6Gi + +# TIA Bridge Config +- name: CPU_TIA_BRIDGE_REQ + value: 1 +- name: CPU_TIA_BRIDGE_LIM + value: 2 +- name: MEM_TIA_BRIDGE_REQ + value: 2Gi +- name: MEM_TIA_BRIDGE_LIM + value: 4Gi + +name: 0.1.$(Rev:r) + +stages: +- stage: Build + jobs: + - job: Build + steps: + - script: | + echo "$(auth-token)" | docker login $(ociRegistry) -u "$(OCI_NAMESPACE)/$(SA_NAME)" --password-stdin + displayName: 'Login to OCI Registry' + + - script: | + set -eo pipefail + + docker build -f $(Build.SourcesDirectory)/k8s/tia/Dockerfile -t $(imageRepository):$(Build.BuildNumber) $(Build.SourcesDirectory) + + docker build -f $(Build.SourcesDirectory)/k8s/livekit/Dockerfile -t $(imageLivekit):$(Build.BuildNumber) $(Build.SourcesDirectory) + displayName: 'Build the image' + + - script: | + docker push $(imageRepository):$(Build.BuildNumber) + docker push $(imageLivekit):$(Build.BuildNumber) + displayName: 'Push image to OCI Registry' + +- stage: Deploy + dependsOn: Build + condition: succeeded() + jobs: + - job: deployb8s + pool: + name: oci-kubernetes + variables: + K8S_NAMESPACE: 'agnt-ai-atendimento-tia' + IMAGE_REPOSITORY: '$(imageRepository)' + IMAGE_REPOSITORY_LIVEKIT: '$(imageLivekit)' + IMAGE_TAG: '$(Build.BuildNumber)' + APP_NAME: '$(Build.Repository.Name)' + OCI_CLI_CONFIG_FILE: '$(Pipeline.Workspace)/oci_config' + + steps: + + - template: oke_agents_atendimento/autenticacao_kubernetes_v2.yaml@templates + + - template: oke_agents_atendimento/apply_env_as_configmap.yaml@templates + parameters: + envFile: '.env.kube.fqa' + + - template: oke_agents_atendimento/sa_gcp_secret.yaml@templates + + - script: | + set -eo pipefail + envsubst < k8s/secrets.yaml | kubectl apply -n $(K8S_NAMESPACE) -f - + displayName: "aplicando secrets" + env: + AZURE_SPEECH_KEY: $(AZURE_SPEECH_KEY) + XAI_API_KEY: $(fXAI_API_KEY_FQA) + LIVEKIT_REDIS_USERNAME: $(REDIS_FQA_USERNAME) + LIVEKIT_REDIS_PASSWORD: $(REDIS_FQA_PASSWORD) + + - script: | + set -eo pipefail + + # cada pipeline que vamos querer deployar + mkdir -p helm_build/templates + + envsubst < k8s/livekit/configmap.yaml > helm_build/templates/configmap-livekit.yaml + envsubst < k8s/livekit/deployment.yaml > helm_build/templates/deployment-livekit.yaml + envsubst < k8s/tia/deployment.yaml > helm_build/templates/deployment-tia.yaml + + envsubst < k8s/helm_deploy.yaml > helm_build/Chart.yaml + helm upgrade --install $(APP_NAME) ./helm_build \ + --namespace $(K8S_NAMESPACE) \ + --wait \ + --timeout 180s \ + --atomic + displayName: 'Apply kube custom helm' \ No newline at end of file diff --git a/azure-pipelines/prod.yml b/azure-pipelines/prod.yml new file mode 100644 index 0000000..d3ec239 --- /dev/null +++ b/azure-pipelines/prod.yml @@ -0,0 +1,156 @@ +trigger: + branches: + include: + - master + +pool: + name: Azure Pipelines + vmImage: 'ubuntu-latest' + +resources: + repositories: + - repository: templates + type: git + name: Data_Science_AI_ML/template-yaml-devops-mlops + ref: master + - repository: TEMPLATES_CICD + type: git + name: TEMPLATES_CICD/TEMPLATES_CICD + ref: main + +variables: +- template: oke_agents_atendimento/variables.yaml@templates +- group: sa-timsa-it-agnt-ai-atendimento-cicd-prd-auth-token +- group: tim-ai-atend-agnt-integ-tia-dev-fqa # parece que prod é a mesma que fqa +- name: imageRepository + value: '$(ociRegistry)/$(OCI_NAMESPACE)/agnt-ai-atendimento/prd/agnt-ai-atendimento-tia' +- name: imageLivekit + value: '$(ociRegistry)/$(OCI_NAMESPACE)/agnt-ai-atendimento/prd/agnt-ai-atendimento-tia-livekit' +- name: DNS + value: agt-ai-atendimento-tia.internal.timbrasil.com.br +- name: REDIS_IP + value: 10.152.89.43 +- name: REDIS_HOST + value: aaaaehl73aawjafeo24e5zuqosggsrnjzlomctu4h636pxdhuf4bspa-p.redis.sa-saopaulo-1.oci.oraclecloud.com +# Livekit Config +- name: LIVEKIT_REPLICAS + value: 3 +- name: CPU_LIVEKIT_REQ + value: 2 +- name: CPU_LIVEKIT_LIM + value: 8 +- name: MEM_LIVEKIT_REQ + value: 2Gi +- name: MEM_LIVEKIT_LIM + value: 12Gi + +# TIA Config +- name: TIA_REPLICAS + value: 10 +# esse valor de réplica vale para ambos os abaixo +- name: CPU_TIA_REQ + value: 1 +- name: CPU_TIA_LIM + value: 6 +- name: MEM_TIA_REQ + value: 4Gi +- name: MEM_TIA_LIM + value: 8Gi + +# TIA Bridge Config +- name: CPU_TIA_BRIDGE_REQ + value: 1 +- name: CPU_TIA_BRIDGE_LIM + value: 4 +- name: MEM_TIA_BRIDGE_REQ + value: 2Gi +- name: MEM_TIA_BRIDGE_LIM + value: 4Gi + +name: 0.0.$(Rev:r) + +stages: +- stage: Build + jobs: + - job: Build + steps: + - script: | + echo "$(auth-token)" | docker login $(ociRegistry) -u "$(OCI_NAMESPACE)/$(SA_NAME)" --password-stdin + displayName: 'Login to OCI Registry' + + - script: | + set -eo pipefail + + docker build -f $(Build.SourcesDirectory)/k8s/tia/Dockerfile -t $(imageRepository):$(Build.BuildNumber) $(Build.SourcesDirectory) + + docker build -f $(Build.SourcesDirectory)/k8s/livekit/Dockerfile -t $(imageLivekit):$(Build.BuildNumber) $(Build.SourcesDirectory) + displayName: 'Build the image' + + - script: | + docker push $(imageRepository):$(Build.BuildNumber) + docker push $(imageLivekit):$(Build.BuildNumber) + displayName: 'Push image to OCI Registry' + +- stage: Deploy + dependsOn: Build + condition: succeeded() + jobs: + - deployment: deployb8s + displayName: 'Deploy Application' + environment: 'agents_pnc_producao' + pool: + name: oci-kubernetes + variables: + K8S_NAMESPACE: 'agnt-ai-atendimento-tia' + IMAGE_REPOSITORY: '$(imageRepository)' + IMAGE_REPOSITORY_LIVEKIT: '$(imageLivekit)' + IMAGE_TAG: '$(Build.BuildNumber)' + APP_NAME: '$(Build.Repository.Name)' + OCI_CLI_CONFIG_FILE: '$(Pipeline.Workspace)/oci_config' + strategy: + runOnce: + deploy: + steps: + - checkout: self + + - template: oke_agents_atendimento/autenticacao_kubernetes_v2.yaml@templates + + - template: oke_agents_atendimento/apply_env_as_configmap.yaml@templates + parameters: + envFile: '.env.kube.prd' + + - template: oke_agents_atendimento/sa_gcp_secret.yaml@templates + + - script: | + set -eo pipefail + envsubst < k8s/secrets.yaml | kubectl apply -n $(K8S_NAMESPACE) -f - + displayName: "aplicando secrets" + env: + AZURE_SPEECH_KEY: $(AZURE_SPEECH_KEY) + XAI_API_KEY: $(fXAI_API_KEY_PRD) + LIVEKIT_REDIS_USERNAME: $(REDIS_PRD_USERNAME) + LIVEKIT_REDIS_PASSWORD: $(REDIS_PRD_PASSWORD) + + - script: | + set -eo pipefail + + mkdir -p helm_build/templates + envsubst < k8s/livekit/configmap.yaml > helm_build/templates/configmap-livekit.yaml + envsubst < k8s/livekit/deployment.yaml > helm_build/templates/deployment-livekit.yaml + envsubst < k8s/tia/deployment.yaml > helm_build/templates/deployment-tia.yaml + + envsubst < k8s/helm_deploy.yaml > helm_build/Chart.yaml + helm upgrade --install $(APP_NAME) ./helm_build \ + --namespace $(K8S_NAMESPACE) \ + --wait \ + --timeout 180s \ + --atomic + displayName: 'Apply kube custom helm' + +# - stage: CreateWIT +# displayName: 'Create WIT Entrega (Manual)' +# dependsOn: Deploy +# jobs: +# - template: wi_entrega/wi_entrega.yaml@templates +# parameters: +# environmentName: 'WI_entrega' \ No newline at end of file diff --git a/ddl/README.md b/ddl/README.md new file mode 100644 index 0000000..c1aa672 --- /dev/null +++ b/ddl/README.md @@ -0,0 +1,3 @@ +# Introduction + DDLs para criação de tabelas ou views + \ No newline at end of file diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..e319341 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,150 @@ +# Stack local de observabilidade com Langfuse. +# A aplicacao principal continua sendo executada via `make agent`, `make bridge` e `make livekit`. + +x-langfuse-env: &langfuse-env + NEXTAUTH_URL: http://localhost:3005 + DATABASE_URL: postgresql://postgres:postgres@postgres:5432/postgres + POSTGRES_USER: postgres + POSTGRES_PASSWORD: postgres + POSTGRES_DB: postgres + SALT: devsalt + ENCRYPTION_KEY: b127cbb367ba27ddf3851750686b88a984acd818c0b8444e9370d11fb75fb7df + NEXTAUTH_SECRET: devsecret + TELEMETRY_ENABLED: "false" + LANGFUSE_ENABLE_EXPERIMENTAL_FEATURES: "true" + CLICKHOUSE_MIGRATION_URL: clickhouse://clickhouse:9000 + CLICKHOUSE_URL: http://clickhouse:8123 + CLICKHOUSE_USER: clickhouse + CLICKHOUSE_PASSWORD: clickhouse + CLICKHOUSE_CLUSTER_ENABLED: "false" + REDIS_HOST: redis + REDIS_PORT: 6379 + REDIS_AUTH: devredis + REDIS_TLS_ENABLED: "false" + LANGFUSE_USE_AZURE_BLOB: "false" + LANGFUSE_S3_EVENT_UPLOAD_BUCKET: langfuse + LANGFUSE_S3_EVENT_UPLOAD_REGION: auto + LANGFUSE_S3_EVENT_UPLOAD_ACCESS_KEY_ID: minio + LANGFUSE_S3_EVENT_UPLOAD_SECRET_ACCESS_KEY: miniosecret + LANGFUSE_S3_EVENT_UPLOAD_ENDPOINT: http://minio:9000 + LANGFUSE_S3_EVENT_UPLOAD_FORCE_PATH_STYLE: "true" + LANGFUSE_S3_EVENT_UPLOAD_PREFIX: events/ + LANGFUSE_S3_MEDIA_UPLOAD_BUCKET: langfuse + LANGFUSE_S3_MEDIA_UPLOAD_REGION: auto + LANGFUSE_S3_MEDIA_UPLOAD_ACCESS_KEY_ID: minio + LANGFUSE_S3_MEDIA_UPLOAD_SECRET_ACCESS_KEY: miniosecret + LANGFUSE_S3_MEDIA_UPLOAD_ENDPOINT: http://minio:9000 + LANGFUSE_S3_MEDIA_UPLOAD_FORCE_PATH_STYLE: "true" + LANGFUSE_S3_MEDIA_UPLOAD_PREFIX: media/ + LANGFUSE_S3_BATCH_EXPORT_ENABLED: "false" + LANGFUSE_S3_BATCH_EXPORT_BUCKET: langfuse + LANGFUSE_S3_BATCH_EXPORT_PREFIX: exports/ + LANGFUSE_S3_BATCH_EXPORT_REGION: auto + LANGFUSE_S3_BATCH_EXPORT_ENDPOINT: http://minio:9000 + LANGFUSE_S3_BATCH_EXPORT_EXTERNAL_ENDPOINT: http://localhost:9090 + LANGFUSE_S3_BATCH_EXPORT_ACCESS_KEY_ID: minio + LANGFUSE_S3_BATCH_EXPORT_SECRET_ACCESS_KEY: miniosecret + LANGFUSE_S3_BATCH_EXPORT_FORCE_PATH_STYLE: "true" + +x-langfuse-depends-on: &langfuse-depends-on + postgres: + condition: service_healthy + redis: + condition: service_healthy + clickhouse: + condition: service_healthy + minio: + condition: service_started + +services: + langfuse-web: + image: docker.io/langfuse/langfuse:3 + restart: unless-stopped + depends_on: *langfuse-depends-on + ports: + - "3005:3000" + environment: + <<: *langfuse-env + + langfuse-worker: + image: docker.io/langfuse/langfuse-worker:3 + restart: unless-stopped + depends_on: *langfuse-depends-on + environment: + <<: *langfuse-env + + postgres: + image: docker.io/postgres:17 + restart: unless-stopped + environment: + POSTGRES_USER: postgres + POSTGRES_PASSWORD: postgres + POSTGRES_DB: postgres + TZ: UTC + PGTZ: UTC + ports: + - "127.0.0.1:5433:5432" + volumes: + - langfuse_postgres_data:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U postgres"] + interval: 3s + timeout: 3s + retries: 20 + + redis: + image: docker.io/redis:7 + restart: unless-stopped + command: + - --requirepass + - devredis + - --maxmemory-policy + - noeviction + ports: + - "127.0.0.1:6379:6379" + healthcheck: + test: ["CMD", "redis-cli", "-a", "devredis", "ping"] + interval: 3s + timeout: 10s + retries: 20 + + clickhouse: + image: docker.io/clickhouse/clickhouse-server + restart: unless-stopped + user: "101:101" + environment: + CLICKHOUSE_DB: default + CLICKHOUSE_USER: clickhouse + CLICKHOUSE_PASSWORD: clickhouse + ports: + - "127.0.0.1:8124:8123" + - "127.0.0.1:9002:9000" + volumes: + - langfuse_clickhouse_data:/var/lib/clickhouse + - langfuse_clickhouse_logs:/var/log/clickhouse-server + healthcheck: + test: ["CMD-SHELL", "wget --no-verbose --tries=1 --spider http://localhost:8123/ping || exit 1"] + interval: 5s + timeout: 5s + retries: 20 + start_period: 5s + + minio: + image: minio/minio + restart: unless-stopped + entrypoint: sh + command: -c 'mkdir -p /data/langfuse && minio server --address ":9000" --console-address ":9001" /data' + environment: + MINIO_ROOT_USER: minio + MINIO_ROOT_PASSWORD: miniosecret + ports: + - "9090:9000" + - "127.0.0.1:9091:9001" + volumes: + - langfuse_minio_data:/data + +volumes: + langfuse_postgres_data: + langfuse_clickhouse_data: + langfuse_clickhouse_logs: + langfuse_minio_data: diff --git a/docs/README.md b/docs/README.md new file mode 100644 index 0000000..b018f08 --- /dev/null +++ b/docs/README.md @@ -0,0 +1,17 @@ +# Docs + +Esta pasta concentra a documentacao viva do projeto. + +Arquivos: +- `../README.md`: setup local, bootstrap do virtualenv, faixa de Python suportada, comandos de teste e estrutura de imagens/manifests Kubernetes +- `refactor-plan.md`: alvo arquitetural e fases do refactor incremental +- `refactor-log.md`: historico operacional das mudancas feitas no repositorio +- `api-overview.md`: descricao da API, fluxos e responsabilidades atuais + +Regra de uso durante o refactor: +1. toda mudanca relevante de arquitetura deve atualizar `refactor-log.md` +2. quando o comportamento publico mudar, atualizar `api-overview.md` +3. se a ordem das etapas do refactor mudar, atualizar `refactor-plan.md` + +Objetivo final: +- chegar ao fim do refactor com uma base documental suficiente para gerar a documentacao final da API sem depender de reconstruir contexto depois diff --git a/docs/api-overview.md b/docs/api-overview.md new file mode 100644 index 0000000..bf7bffb --- /dev/null +++ b/docs/api-overview.md @@ -0,0 +1,728 @@ +# API Overview + +## Status + +Documento vivo. Descreve o estado atual da API e sera atualizado ao longo do refactor. + +## O que esta API faz + +Esta API implementa um fluxo de voz-to-voz para TIA: +- recebe conexao websocket do cliente +- recebe audio PCM do cliente +- encaminha audio para o agent via LiveKit +- transcreve, processa a pipeline de negocio e sintetiza resposta +- devolve audio para o cliente +- finaliza a chamada e exporta resultado + +## Componentes principais + +### Bridge + +Arquivo principal: +- `app/ws_gateway/main.py` +- `app/ws_gateway/voice_client.html` + +Responsabilidades: +- aceitar conexao em `/ws/agent` +- receber `start` +- montar contexto da chamada +- despachar agent para uma room LiveKit +- mandar audio do cliente para LiveKit +- devolver audio do agent para o cliente +- encerrar websocket ao receber `DONE` +- servir um cliente web de teste em `/voice-client` + +### Agent + +Arquivos principais: +- `app/livekit/main.py` +- `app/livekit/runtime/call_runtime.py` +- `app/livekit/runtime/state.py` +- `app/livekit/runtime/commands.py` +- `app/livekit/runtime/command_executor.py` +- `app/livekit/runtime/scheduler.py` +- `app/livekit/policies/` +- `app/livekit/adapters/` + +Responsabilidades: +- entrar na room do LiveKit +- configurar STT, VAD e TTS +- receber transcricao final do usuario +- chamar o backend de IA configurado +- vocalizar resposta +- tratar interrupcao, idle nudge e finalizacao + +Organizacao atual: +- `main.py` faz o wiring do agent e das dependencias +- `CallRuntime` coordena o ciclo de vida da chamada +- `CallState` concentra o estado mutavel da sessao +- `commands.py` define os comandos internos para side effects +- `RuntimeCommandExecutor` executa bridge, export, speech, pipeline e start da sessao +- `TimerScheduler` concentra os timers nomeados do runtime +- `InterruptPolicy`, `IdlePolicy` e `FinalizationPolicy` concentram regras operacionais da chamada +- adapters encapsulam acesso ao backend de IA, speech, bridge e export + +### Pipeline de negocio + +Backends disponiveis no runtime websocket: +- `remote_ws`: + - `app/livekit/adapters/remote_agent_ws_adapter.py` +- `remote_sse`: + - `app/livekit/adapters/remote_agent_sse_adapter.py` + +Responsabilidades: +- preparar dados de atendimento +- controlar os estagios da chamada ou delegar esse controle ao agent remoto +- gerar a resposta textual por etapa +- registrar metadados e finalizar a conversa + +Selecao atual: +- `AGENT_BACKEND=remote_ws` envia cada turno transcrito para `REMOTE_AGENT_WS_URL` +- `AGENT_BACKEND=remote_sse` usa contratos especificos por agente: + - `conta` usa `GET /agent/sse` para inicializar a sessao e `POST /agent/sse` para executar cada acao + - `oferta` usa `POST /agent/execute` por turno e consome os eventos SSE `schedule_message`, `message` e `done` +- `AGENT_BACKEND=remote_ws_fake` usa um fake interno no proprio processo, sem abrir websocket local +- o fluxo do runtime local continua o mesmo: STT produz texto, o backend de IA devolve texto e o TTS vocaliza +- o roteamento entre agentes acontece pelo campo `data.agent` +- quando `AGENT_BACKEND=remote_ws`, o adapter escolhe a URL por `agent` se existirem: + - `REMOTE_AGENT_WS_URL_CONTA` + - `REMOTE_AGENT_WS_URL_OFERTA` + - `REMOTE_AGENT_WS_URL_COBRANCA` ou `REMOTE_AGENT_WS_URL_COBRA` +- quando `AGENT_BACKEND=remote_sse`, o adapter exige a URL especifica do `agent`, sem fallback generico: + - `REMOTE_AGENT_SSE_URL_CONTA` + - `REMOTE_AGENT_SSE_URL_OFERTA` + - `REMOTE_AGENT_SSE_URL_COBRANCA` ou `REMOTE_AGENT_SSE_URL_COBRA` + +Contrato remoto atual por turno: +- `timestamp` +- `agent` +- `RouterCallKeyDay` +- `RouterCallKey` +- `ANI` +- `GSM` +- `callIdGed` +- `ID_FATURA` somente para `agent=conta` +- `text` +- `protocol` +- `stage` + +Contrato especifico atual de `conta`: +- request: + - na abertura da sessao, o transporte envia query string com `ani`, `channelId` e `uraCallId` + - `action: "chat"` + - `payload.message` + - `payload.message_id` no transporte SSE de `conta`, com o mesmo UUID enviado na query string + - `payload.channel` no transporte websocket + - `payload.interruption` e `payload.events` quando existirem +- response: + - `type: "ready"` para a primeira fala ou fala que abre uma janela de resposta do cliente + - a fala de `ready` nao e interrompivel; se o cliente falar enquanto o audio do `ready` ainda estiver tocando, a transcricao e descartada e registrada em log + - apos o fim do audio de `ready`, o runtime mantem uma janela protegida de 750ms para absorver atraso de playback do cliente; fala iniciada nessa janela tambem e descartada + - depois dessa janela protegida, o runtime passa a esperar resposta do cliente + - `type: "result"` + - `action: "chat"` + - `result.type: "final"` para respostas intermediarias + - `result.content` como texto para TTS + - `type: "feedback"` ou `result.type: "feedback"` para mensagens de acompanhamento enquanto o backend continua processando o turno atual + - a fala de `feedback` nao e interrompivel, nao abre novo turno, nao espera resposta do cliente e nao arma timeout de silencio do cliente + - se o cliente falar durante `feedback`, o runtime descarta a transcricao, registra a tentativa em log e continua aguardando a resposta final do backend + - `feedback` nao dispara `metadata.wait_retry_messages`, como mensagens do tipo "Voce esta ai?" + - para finalizacoes esperadas, o runtime aguarda uma janela curta de silencio estavel do cliente, fala o `result.content` e depois envia `stop`; se uma nova transcricao final chegar durante a espera, a finalizacao anterior e descartada em favor do novo turno + - falas de finalizacao esperada (`resolvido`, `nao_resolvido`, `erro_no_match` etc.) nao sao interrompiveis; qualquer fala do cliente durante a finalizacao e descartada + - nesses casos o runtime nao chama finalizacao remota adicional (`end`/`end_service_once`) + - quando o `ready` vier com `metadata.wait_timeout_seconds`, o runtime aguarda esse tempo apos falar a mensagem de `ready` + - se tambem vier `metadata.wait_retry_messages`, o runtime fala cada item do array a cada novo estouro de `wait_timeout_seconds`; depois do ultimo item, aguarda mais um intervalo igual e envia + `stop_silencio_longo` com `reason: "no_user_response"` + - sem `metadata.wait_retry_messages`, o comportamento continua sendo encerrar direto no primeiro estouro de `wait_timeout_seconds` + - enquanto aguarda processamento depois do fim de fala do usuario, o runtime toca somente o audio local longo de conforto; no padrao atual, usa intervalo de 12s e no maximo 6 vezes + - se a resposta do backend remoto nao chegar em 300s, o runtime encerra com `stop_agent_backend_unavailable` +- se o participante do agent LiveKit desconectar sem `DONE`, o bridge tenta redispatch na mesma room; se o agent nao voltar, encerra com `stop_agent_runtime_unavailable` + +Mapeamento de finalizacoes esperadas de `conta`: +- `result.type: "resolvido"` -> `stop_resolvido_e_finalizado` +- `result.type: "nao_resolvido"` -> `stop_nao_resolvido` +- `result.type: "resolvido_outros_assuntos"` -> `stop_outro_assunto` +- `result.type: "outros_assuntos"` -> `stop_outro_assunto` +- `result.type: "erro_falha_sistema"` -> `stop_falha_sistema` +- `result.type: "erro_no_match"` -> `stop_no_match` + +Contrato SSE atual de `conta`: +- o runtime guarda o `session_id` retornado pelo backend nos eventos `ready` +- abre `GET /agent/sse?msisdn=...&invoice_id=...&ani=...&protocol_id=...&session_id=...&message_id=...&channelId=ura&uraCallId=...` no `prepare` +- consome o stream de inicializacao ate `ready` e ate o termino do prefetch (`prefetch_done`, `prefetch_skipped` ou `prefetch_failed`) ou fechamento da resposta +- envia cada turno com `POST /agent/sse?session_id=...&ani=...&protocol_id=...&message_id=...&channelId=ura&uraCallId=...` +- `message_id` e um UUID gerado por turno e independente de `session_id`; nos turnos de acao, `session_id` continua sendo o identificador de sessao retornado pelo backend de contas +- body do `POST`: + - `action: "chat"` + - `payload.message` + - `payload.message_id` com o mesmo UUID do parametro `message_id` da query string + - `payload.interruption` e `payload.events` quando existirem +- a resposta do proprio `POST` e `text/event-stream`; o runtime consome `ready`, `progress`, `result` e `error` ate receber o resultado da acao + +Contrato SSE atual de `oferta`: +- `prepare` nao abre stream remoto; o primeiro turno do agente chama `POST /agent/execute` +- o turno inicial envia `message: "inicio_atendimento"` para o backend remoto produzir a primeira fala +- cada turno usa body: + - `messageId` gerado por turno como UUID + - `message` com `inicio_atendimento` no primeiro turno ou a transcricao do cliente nos demais + - `context.protocolNumber` e `context.protocolo` vindos de `data.protocolo` + - `context.gsm` vindo de `data.gsm` + - `context.uraId` vindo de `data.callIdGed` + - `context.callIdGed`, `context.ani`, `context.routerCallKey`, `context.routerCallKeyDay`, `context.agent` e `context.assetId` quando existirem +- `context.sessionId` e enviado a partir de `data.session_id`/`data.sessionId` recebido no `start`; protocolo, `callIdGed` e `routerCallKey` nao sao usados como fallback de `sessionId` +- headers enviados: `Accept: text/event-stream`, `Content-Type: application/json` e `Channel-id` vindo de `channelId` ou `ura` +- eventos recebidos: + - `schedule_message`: fala imediatamente `scheduledMessage.message` + - `message`: fala `response` + - `done`: encerra o stream do turno; a chamada so e finalizada quando o status recebido indicar encerramento terminal +- de/para do `done.additionalInformations.service_status`: + - `RESOLVED` -> `stop_resolvido_e_finalizado` + - `UNRESOLVED` ou ausente -> `stop_nao_resolvido` + - `RESOLVED_WITH_NEW_REQUEST` -> nao envia `stop`; a conversa permanece aberta para o novo assunto +- `done.status=transferred` e reconhecido, mas ainda nao orquestra transferencia; nesta versao encerra como nao resolvido e preserva metadados de handover para evolucao futura + +Endpoint dev fornecido para `oferta`: +- host: `https://agt-ai-atendimento-ofertas-dev.internal.timbrasil.com.br` +- execute: `https://agt-ai-atendimento-ofertas-dev.internal.timbrasil.com.br/agent/execute` +- health: nao configurado por enquanto; a readiness nao deve derivar nem chamar `/health` para oferta ate essa rota ser fornecida +- enquanto o certificado interno nao estiver confiavel no ambiente local, `REMOTE_AGENT_SSE_TLS_VERIFY_OFERTA=0` permite testar o fluxo ignorando a validacao TLS do `httpx` +- pod de referencia: `tim-ai-atend-agnt-sales-65764fcf8d-xpzmb` +- IP de referencia: `http://10.153.35.23` +- portas: `80:31332/TCP`, `443:30635/TCP` + +Origem desses campos: +- o bridge extrai esses valores do `data` recebido em `WS /ws/agent` +- o agent local apenas reaproveita esse contexto para chamar o websocket remoto +- `protocol`/`protocolo` e identificador de negocio; `session_id` e identificador explicito de sessao recebido no `start`. O runtime nao preenche `session_id` com protocolo, `callIdGed` ou `routerCallKey`. + +Configuracao opcional por chamada: +- o cliente pode enviar `callConfig` no `start`, mas esse bloco e opcional +- objetivo atual: testes, homologacao e overrides tecnicos por sessao +- `callConfig.agentBackend` faz override do backend websocket para aquela chamada +- `callConfig.agentBackend` tambem pode selecionar o transporte SSE para a sessao +- `callConfig.stt` pode ajustar `provider`, `initialPrompt`, `configOverride` e `minProbSingleWord` +- `callConfig.tts` pode ajustar `provider`, `voiceId` e `modelId` +- valores aceitos hoje em `callConfig.agentBackend`: + - `remote_ws` + - `remote_sse` + - `remote_ws_fake` +- hoje os providers suportados no agent local sao: + - STT: `internal_http` (`Sofya Batch` no cliente de teste), `fake` + - TTS: `elevenlabs`, `azure`, `xai`, `fake` + +## Endpoints atuais + +### `GET /health` + +Retorna status simples de saude do bridge. + +### `GET /health/resources` + +Retorna o deep health do bridge com validacao dos recursos usados pelo `WS /ws/agent`. + +Comportamento: +- `200` quando os checks obrigatorios estao saudaveis +- `503` quando algum recurso obrigatorio falha +- inclui `active_connections`, `max_connections`, `cached`, `failed_resources` e `checks` +- os checks atuais cobrem: + - `agent_runtime` + - `agent_backend` + - `stt` + - `tts` + +### `GET /health/services` + +Retorna o health consolidado dos servicos do TIA no formato consumido pela esteira de operacao. + +Comportamento: +- `200` quando todos os servicos monitorados estao saudaveis +- `503` quando algum servico monitorado falha +- `status` no corpo retorna `ok` ou `fail` +- `htp_cod_status` preserva o status HTTP retornado pelo health do servico quando houver resposta HTTP +- `hhtp_cod_desc` preserva a descricao retornada pelo health quando houver; se nao houver descricao no corpo, usa a reason phrase HTTP +- ignora `AGENT_BACKEND=remote_ws_fake` e usa as rotas reais configuradas nos envs dos servicos +- `checks` cobre apenas os servicos atualmente monitorados: + - `agent_runtime` + - `agent_backend.contas`: `REMOTE_AGENT_HEALTH_URL_CONTA`, `REMOTE_AGENT_HEALTH_URL_CONTAS` ou `REMOTE_AGENT_HEALTH_URL` + - `agent_backend.oferta`: `REMOTE_AGENT_HEALTH_URL_OFERTA`, `REMOTE_AGENT_HEALTH_URL_OFERTAS` ou health derivado de `REMOTE_AGENT_SSE_URL_OFERTA` + - `stt.sofya`: `STT_HEALTH_URL` ou health derivado de `STT_URL` + - `tts.xAI`: provider xAI com `XAI_WEBSOCKET_URL`/`XAI_API_KEY` + +Para checks sem resposta HTTP, como falha de conexao ou probe por WebSocket, a rota usa fallback `200`/`500` e a mensagem interna do erro ou sucesso. + +Formato: + +```json +{ + "status": "ok", + "checks": { + "agent_runtime": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS" + }, + "agent_backend": { + "contas": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS" + }, + "oferta": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS" + } + }, + "stt": { + "sofya": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS" + } + }, + "tts": { + "xAI": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS" + } + } + } +} +``` + +### `GET /voice-client` + +Cliente web de teste para capturar microfone, enviar audio para `WS /ws/agent`, +reproduzir o audio do agent e configurar `STT`, `TTS` e `AGENT` por chamada. + +### `WS /fake-agent/ws` + +Websocket fake para homologacao local do backend remoto. + +Uso esperado: +- usar apenas para homologacao manual do contrato websocket fake +- manter `agent=conta|oferta|cobranca` +- deixar um cliente websocket externo chamar o fake diretamente quando precisar testar esse endpoint + +Comportamento: +- suporta o contrato `conta` com `action/payload` +- suporta o contrato generico com `text/stage` +- responde com progressao simples de stages +- encerra quando recebe textos como `encerrar`, `obrigado` ou `tchau` + +### `WS /ws/agent` + +Fluxo principal de voz: +1. cliente conecta +2. envia mensagem `start` +3. bridge valida capacidade e readiness dos recursos obrigatorios +4. recebe `ready` ou `stop` +5. se receber `ready`, envia audio binario +6. recebe audio binario de resposta +7. recebe `stop` ao fim da chamada ou em falha terminal + +Contrato de inicio da chamada `type=start`: +- a primeira mensagem de inicio deve ser um JSON textual com `type: "start"`; opcionalmente, `transferencia_session_id` pode chegar antes dela para informar somente o `session_id` +- os campos de negocio devem ser enviados em `data` +- chaves canonicas obrigatorias em `data`: + - `agent` + - `ani` + - `gsm` + - `session_id` + - `routerCallKey` + - `routerCallKeyDay` + - `callIdGed` +- chaves obrigatorias por agente: + - `agentData.idFatura` quando `agent=conta` + - `protocolo` quando `agent=oferta` +- valores canonicos recomendados para `agent`: + - `conta` + - `oferta` + - `cobranca` + +Exemplo recomendado de `start` para `conta`: + +```json +{ + "type": "start", + "data": { + "agent": "conta", + "ani": "5511999990000", + "gsm": "5511999990000", + "session_id": "550e8400-e29b-41d4-a716-446655440000", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-123456", + "agentData": { + "idFatura": "FAT-123" + } + }, + "audioFormat": { + "encoding": "linear16", + "sampleRateHz": 16000, + "channels": 1 + }, + "callConfig": { + "agentBackend": "remote_ws" + } +} +``` + +Exemplo recomendado de `start` para `oferta`: + +```json +{ + "type": "start", + "data": { + "agent": "oferta", + "ani": "5511999990000", + "gsm": "5511999990000", + "session_id": "550e8400-e29b-41d4-a716-446655440000", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-123456", + "protocolo": "PRT-20260409-0001" + }, + "audioFormat": { + "encoding": "linear16", + "sampleRateHz": 16000, + "channels": 1 + }, + "callConfig": { + "agentBackend": "remote_sse" + } +} +``` + +Mensagem opcional aceita antes do `start` em cenarios de transferencia: + +```json +{ + "type": "transferencia_session_id", + "data": { + "session_id": "550e8400-e29b-41d4-a716-446655440000" + } +} +``` + +Se essa mensagem chegar antes do `start`, o bridge usa esse valor apenas para preencher `data.session_id` quando o `start` ainda nao trouxer o campo. Em producao, o caminho recomendado e enviar `session_id` diretamente dentro de `data` no `start`. + +Contrato de audio: +- apos o `ready`, o cliente deve enviar audio binario bruto em `PCM16/LINEAR16`, `16000 Hz`, `1 canal` +- o `ready` informa os parametros operacionais atuais do bridge: + - `session_id`: eco do `data.session_id` aceito para a chamada + - `sample_rate: 16000` + - `channels: 1` + - `frame_ms: 20` + - `bytes_per_frame: 640` +- no contrato atual, `ready` significa que o bridge esta pronto para receber os bytes de audio do cliente +- quando `agent_starts_conversation` esta habilitado, o agent mantem a entrada do usuario desativada no `RoomIO` desde antes do `StartSession` ate o fim da primeira mensagem; esse gate impede que fala ou backlog vindo da URA alcance VAD/STT durante o setup e a saudacao +- o gate e liberado no fim do primeiro turno do agent e tambem em falha do pipeline, finalizacao ou timeout de setup; chamadas em que o usuario inicia a conversa nao usam esse bloqueio +- `audioFormat` no `start` e opcional e hoje funciona como campo informativo/reservado +- enviar outro codec, sample rate ou numero de canais nesse campo nao reconfigura o bridge atualmente +- se houver necessidade de outro formato, homologar com a equipe de desenvolvimento antes da integracao + +Mensagens devolvidas pelo servidor: +- `ready` quando a sessao foi aceita e o bridge esta pronto para receber audio +- audio binario PCM16 durante a resposta do agent +- `stop` como mensagem terminal em qualquer encerramento do `WS /ws/agent` + +Recuperacao do participante LiveKit do agent: +- quando o participant identificado como agent desconecta e a chamada ainda nao terminou, o bridge faz redispatch do mesmo `AGENT_NAME` na mesma room +- quando o novo participant entra, o bridge reenvia o controle `client_audio_enabled` e passa a consumir o audio desse novo participant +- se o redispatch nao trouxer um novo agent dentro do timeout configurado, o bridge envia `stop_agent_runtime_unavailable` com `reason: "agent_disconnected"` + +Contrato de `stop`: +- toda mensagem terminal em `WS /ws/agent` usa: + +```json +{ + "type": "stop", + "data": {} +} +``` + +- bloqueio antes do `ready`: + +```json +{ + "type": "stop", + "data": { + "status": "stop_stt_unavailable", + "reason": "resource_unhealthy", + "resource": "stt", + "failed_resources": ["stt"], + "phase": "pre_ready" + } +} +``` + +- falha de recurso durante a sessao: + +```json +{ + "type": "stop", + "data": { + "status": "stop_agent_backend_unavailable", + "reason": "resource_unhealthy", + "resource": "agent_backend", + "failed_resources": ["agent_backend"], + "phase": "in_session" + } +} +``` + +- fim normal da chamada: + +```json +{ + "type": "stop", + "data": { + "status": "stop_resolvido_e_finalizado", + "reason": "stage_done", + "phase": "in_session" + } +} +``` + +- falha terminal durante a sessao: + +```json +{ + "type": "stop", + "data": { + "status": "stop_bridge_failed", + "reason": "bridge_failed", + "resource": "bridge", + "phase": "in_session" + } +} +``` + +Status terminais atualmente usados em `WS /ws/agent`: +- `stop_capacity_tia` +- `stop_agent_runtime_unavailable` +- `stop_agent_backend_unavailable` +- `stop_stt_unavailable` +- `stop_tts_unavailable` +- `stop_resolvido_e_finalizado` +- `stop_nao_resolvido` +- `stop_falha_sistema` +- `stop_no_match` +- `stop_outro_assunto` +- `stop_silencio_longo` +- `stop_bridge_failed` + +Regra de interpretacao: +- `phase: "pre_ready"` indica bloqueio antes de a sessao aceitar audio +- `phase: "in_session"` indica falha terminal ou encerramento depois do `ready` +- os status de recurso podem aparecer nas duas fases, dependendo de quando a indisponibilidade foi detectada + +Campos opcionais de testes e homologacao: +- `audioFormat` pode ser enviado no `start`, mas hoje nao altera o pipeline de audio +- `callConfig` e opcional e existe para testes, smoke test local e homologacao tecnica +- para integracao produtiva com cliente externo, o contrato pode omitir `callConfig` + +Configuracao de chamada para teste de carga com STT Sofya, agente fake e TTS xAI: + +```json +{ + "debugEvents": true, + "callConfig": { + "agentBackend": "remote_ws_fake", + "agentFake": { + "delayMs": 2500, + "responses": "Primeira resposta simulada com tamanho intermediario;Segunda resposta simulada com tamanho intermediario;Resposta final encerrando o atendimento simulado" + }, + "stt": { + "provider": "internal_http", + "disableVosk": true + }, + "tts": { + "provider": "xai" + } + } +} +``` + +Regras desse modo: +- `agentFake.responses` contem de 2 a 10 frases separadas por `;`; espacos laterais sao removidos e cada frase deve ter de 40 a 180 caracteres +- `agentFake.delayMs` e aplicado antes de cada resposta, usa `2500` por padrao e aceita valores de `0` a `180000` +- a saudacao inicial continua sendo o `intro` normal; cada fala posterior reconhecida pelo STT consome uma resposta fake, sem usar o texto transcrito para escolher a resposta +- respostas anteriores as duas ultimas usam `ARGUMENTATION`, a penultima usa `FORMALIZATION` e a ultima usa `DONE`; chamadas posteriores ao fim recebem novamente o mesmo resultado terminal +- a sequencia e isolada por sessao +- `debugEvents=true` publica pelo websocket somente eventos tecnicos e valores agregados de fala, STT Sofya, TTS xAI e sheds do Bridge; o conteudo integral da conversa nao e incluido nas metricas +- o endpoint nao possui uma autorizacao adicional especifica para o fake: qualquer cliente ja autorizado a abrir `/ws/agent` pode selecionar `remote_ws_fake` por `callConfig` + +### `WS /ws/text` + +Fluxo de texto sem audio para testes e integracao basica. + +### `WS /ws/text_stream` + +Fluxo textual com resposta em streaming. + +## Fluxo ponta a ponta atual + +1. Cliente abre websocket em `/ws/agent` +2. Bridge recebe `start` e faz parse do contexto da chamada +3. Bridge valida capacidade e readiness dos recursos +4. Bridge cria room/token e envia `ready` +5. Cliente envia audio para o bridge +6. Bridge publica audio no LiveKit +7. Agent recebe audio, STT produz texto +8. Agent chama o backend de IA configurado +9. Agent usa TTS para vocalizar resposta +10. Bridge devolve audio ao cliente +11. Agent sinaliza `DONE` ou ocorre erro terminal +12. Bridge envia `stop` e fecha a chamada + +## Dependencias externas relevantes + +- LiveKit +- STT interno HTTP +- Vosk +- ElevenLabs +- modelo LLM da pipeline +- API de fidelizacao + +## Timeline de chamada + +O projeto agora gera uma timeline estruturada por chamada em formato `jsonl`. + +Configuracao: +- `CALL_TIMELINE_ENABLED=1` ativa a escrita da timeline +- `CALL_TIMELINE_CONSOLE=1` replica os eventos tambem no stdout +- `CALL_TIMELINE_DIR=./timeline` define o diretorio dos arquivos +- `CALL_TIMELINE_QUEUE_MAX=10000` limita eventos pendentes para escrita assincrona +- `CALL_LOG_QUEUE_MAX=20000` limita registros pendentes dos arquivos por chamada +- `ASYNC_IO_WARNING_INTERVAL_S=60` limita a frequencia dos avisos de descarte/erro + +A timeline e o arquivo de log por chamada sao gravados por threads de fundo para +nao executar I/O de disco no event loop de audio. Quando uma fila atinge o limite, +o evento e descartado em vez de bloquear o audio e um warning rate-limited registra +o total acumulado. Os limites aceitos ficam entre `1` e `1000000`. + +Os spans estruturados usam `BatchSpanProcessor`: `span.end()` apenas enfileira o +span, enquanto o envio OTLP acontece em lote fora da thread chamadora. O provider +faz flush e shutdown no encerramento normal do processo. + +Fake remoto: +- `remote_ws_fake` usa um fake interno em memoria e nao depende de `REMOTE_AGENT_WS_FAKE_URL` +- o endpoint `ws://127.0.0.1:8000/fake-agent/ws` continua disponivel apenas para testes manuais do contrato websocket + +Logs de websocket remoto: +- o adapter `remote_ws` agora registra no stdout eventos `REMOTE_AGENT_WS_CONNECT_OPEN`, `REMOTE_AGENT_WS_CONNECT_OK`, `REMOTE_AGENT_WS_REQUEST`, `REMOTE_AGENT_WS_RESPONSE` e falhas `*_FAIL` +- os logs incluem `instance` (hostname/pod), `agent`, `url`, `host`, `stage`, `protocol` e metadados do payload para facilitar comparar pods com erro de DNS/host + +Mock de encerramento apos primeiro audio: +- `MOCK_STOP_AFTER_FIRST_AUDIO_ENABLED=1` faz o bridge enviar o `stop` terminal configurado para o `reason` depois que o primeiro audio real do agente for entregue ao cliente e a saida ficar em silencio pelo intervalo configurado +- `MOCK_STOP_AFTER_FIRST_AUDIO_SILENCE_S=0.35` controla quanto tempo de silencio o bridge espera antes de disparar o `stop` +- `MOCK_STOP_AFTER_FIRST_AUDIO_REASON=stage_done` define o `reason` exato enviado no `stop` + +Status terminais configuraveis por env: +- `FINAL_STOP_STATUS_RESOLVED=stop_resolvido_e_finalizado` +- `FINAL_STOP_STATUS_UNRESOLVED=stop_nao_resolvido` +- `FINAL_STOP_STATUS_OTHER_SUBJECT=stop_outro_assunto` +- `FINAL_STOP_STATUS_LONG_SILENCE=stop_silencio_longo` +- `FINAL_STOP_DEFAULT_KIND=resolved` define o fallback quando o `reason` nao bater em nenhum valor configurado +- `FINAL_STOP_REASON_RESOLVED=stage_done` +- `FINAL_STOP_REASON_UNRESOLVED=nao_resolvido` +- `FINAL_STOP_REASON_OTHER_SUBJECT=outro_assunto` +- `FINAL_STOP_REASON_LONG_SILENCE=no_user_response` +- o mapeamento agora usa comparacao exata do `reason`, sem aliases + +Espera de resposta do backend remoto: +- `REMOTE_AGENT_INFLIGHT_WAIT_INTERVAL_S=12` +- `REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR=src/app/livekit/assets/comfort/short` +- `REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR=src/app/livekit/assets/comfort/long` +- o primeiro audio de conforto usa um WAV aleatorio da pasta `short`; a partir do segundo, usa WAVs aleatorios da pasta `long` +- o intervalo dos audios de conforto conta a partir do fim do audio anterior +- `REMOTE_AGENT_INFLIGHT_WAIT_TIMEOUT_S=180` +- `REMOTE_AGENT_INFLIGHT_WAIT_MAX_NOTICES=0` (`0` mantém os confortos sem limite de quantidade até o timeout) +- `REMOTE_AGENT_INFLIGHT_WAIT_TEXT=Um momento, ainda estou consultando para te ajudar.` +- este timeout e tecnico: limita quanto tempo o runtime aguarda o backend remoto processar um turno +- ele e diferente do timeout de resposta do cliente configurado por `metadata.wait_timeout_seconds` +- mensagens `feedback`, tanto top-level quanto `result.type: "feedback"`, usam a espera tecnica do backend, mas nao contam silencio do cliente nem disparam `metadata.wait_retry_messages` +- em um turno normal, depois que o agent termina de falar e o cliente responde, o fim de fala detectado pelo VAD pode antecipar o primeiro conforto curto enquanto o STT conclui a transcricao +- o conforto antecipado do VAD nao e agendado enquanto o agent esta falando; se outra fala do agent ocupar o `say_lock` depois do agendamento, o conforto tambem e abortado para nao sair colado ao fim da mensagem +- se o cliente interromper uma fala interrompivel do agent, o VAD nao enfileira conforto atras dessa fala; o novo pipeline ainda pode usar os confortos normais caso o processamento demore +- se o cliente falar por mais de `800ms` enquanto o backend ja processa outro turno e o agent esta em silencio, o conforto curto especulativo do VAD e suprimido e a transcricao abre uma interrupcao diferida +- cada interrupcao diferida toca `src/app/livekit/assets/comfort/interruption/01.wav`, correspondente a "Ouvi o que voce falou, um instante", e invalida a resposta anterior quando ela chegar +- no pipeline substituto, o conforto dedicado conta como o primeiro aviso ja consumido: o audio `short` normal nao toca logo depois; se o processamento continuar, o proximo aviso permitido e `long` e respeita o intervalo configurado +- o estado de interrupcao e rearmado antes de iniciar o pipeline substituto: cada nova fala valida durante o novo processamento repete o conforto dedicado e substitui novamente a resposta em voo, sem descartar a nova transcricao nem deixar silencio na chamada +- falas de ate `800ms` durante processamento sao tratadas como ruido ou backchannel curto e nao abrem interrupcao diferida +- os logs principais desse fluxo sao `pre_backend_wait_notice_skipped` (`agent_speaking` ou `backend_processing`), `deferred_interruption_accumulated`, `deferred_interruption_dispatched` e os estagios TTS `AGENT_BACKEND_WAIT`/`INTERRUPTION_COMFORT` + +Recuperacao do participant LiveKit do agent: +- `AGENT_RECONNECT_ENABLED=1` +- `AGENT_RECONNECT_MAX_ATTEMPTS=1` +- `AGENT_RECONNECT_TIMEOUT_S=10` + +Protecao de reenvio TTS: +- `TTS_EMPTY_FRAME_RETRY_TIMEOUT_S=3` interrompe uma tentativa de TTS sem primeiro frame de audio ou com gap entre frames no meio da fala apos o timeout e reenvia o mesmo texto uma vez +- antes do reenvio, toca `src/app/livekit/assets/comfort/fails/tts_fail_recovery.wav` +- quando o reenvio tem sucesso, registra um unico `envio msg` com `http_cod_status=200` e `erro_msg=TTS_regerado` + +STT fake: +- `STT_PROVIDER=fake` dispensa `STT_URL` +- usa `FAKE_STT_TRANSCRIPTS` como fila de falas por turno +- `FAKE_STT_MODE=repeat_last|cycle` controla o comportamento ao consumir a fila + +TTS fake: +- `TTS_PROVIDER=fake` dispensa credenciais externas +- gera audio PCM sintetico local para smoke tests do pipeline + +Formato: +- um arquivo por chamada +- nome do arquivo baseado no `room` +- eventos do `bridge` e do `agent` entram no mesmo arquivo +- cada linha contem: + - `ts` + - `t_rel_ms` + - `component` + - `event` + - `protocol` + - `room` + - campos especificos do evento + +Eventos relevantes: +- `bridge`: + - `call_start` + - `ready_sent` + - `dispatch_started` + - `livekit_room_connected` + - `agent_join` + - `client_audio_first_frame_received` + - `client_audio_first_frame_published` + - `client_audio_enabled` + - `done_packet_received` + - `stop_sent` + - `call_end` +- `agent`: + - `call_start` + - `room_enter` + - `session_start_requested` + - `stt_recognize_started` + - `stt_http_completed` + - `user_transcript_final` + - `pipeline_run_started` + - `pipeline_run_completed` + - `remote_agent_request` + - `remote_agent_response` + - `tts_stage_started` + - `tts_stage_result` + - `interrupt_marked` + - `finalize_started` + - `finalize_completed` + +Uso pratico: +1. iniciar a chamada normalmente +2. identificar no terminal o `room` ou o `protocol` +3. abrir o arquivo correspondente em `./timeline` +4. ler os eventos em ordem de `t_rel_ms` + +## Ponto de atencao + +Durante o refactor, este documento deve continuar descrevendo: +- comportamento publico +- contrato do websocket +- responsabilidades de cada camada + +Na versao final, ele deve evoluir para a documentacao oficial da API. diff --git a/docs/refactor-log.md b/docs/refactor-log.md new file mode 100644 index 0000000..00884c6 --- /dev/null +++ b/docs/refactor-log.md @@ -0,0 +1,612 @@ +# Refactor Log + +Este arquivo registra o que mudou a cada etapa do refactor incremental. + +Nota: +- as entradas `000` a `010` foram escritas antes da consolidacao da estrutura atual do repositorio +- por isso, varias delas ainda referenciam caminhos antigos sem o prefixo `src/` +- a partir da entrada `011`, os caminhos refletem o layout atual da raiz + `src/` + +## Como preencher + +Para cada mudanca relevante, registrar: +- data +- objetivo +- arquivos alterados +- comportamento preservado +- comportamento alterado +- risco conhecido +- validacao executada +- proximo passo + +## Entrada 000 - Baseline documental + +- Data: 2026-03-27 +- Objetivo: criar uma base de documentacao viva para acompanhar o refactor incremental +- Arquivos alterados: + - `README.md` + - `docs/README.md` + - `docs/refactor-plan.md` + - `docs/refactor-log.md` + - `docs/api-overview.md` +- Comportamento preservado: nenhum codigo de execucao foi alterado +- Comportamento alterado: nenhum +- Risco conhecido: a documentacao precisa ser mantida junto das mudancas, senao perde valor +- Validacao executada: revisao manual do conteudo criado +- Proximo passo: iniciar Fase 1 com extracao dos adapters de pipeline, speech, bridge e export + +## Entrada 001 - Fase 1 / adapters extraidos + +- Data: 2026-03-27 +- Objetivo: extrair boundaries de pipeline, speech, bridge e export sem mudar o comportamento publico da API +- Arquivos alterados: + - `app/livekit/main.py` + - `app/livekit/adapters/pipeline_adapter.py` + - `app/livekit/adapters/speech_service.py` + - `app/livekit/adapters/bridge_gateway.py` + - `app/livekit/adapters/export_service.py` + - `docs/refactor-log.md` +- Comportamento preservado: + - contrato do websocket do bridge + - fluxo STT -> pipeline -> TTS + - notificacao `DONE` para o bridge + - exportacao final de sessao +- Comportamento alterado: + - nenhum comportamento publico planejado + - `app/livekit/main.py` passa a delegar responsabilidades para adapters +- Risco conhecido: + - como a extracao preserva a logica inline, ainda existe acoplamento forte no runtime + - faltam testes automatizados de regressao para interrupcao, idle nudge e finalizacao +- Validacao executada: + - revisao manual do diff + - validacao sintatica prevista apos a extracao +- Proximo passo: + - reduzir `nonlocal` no runtime com um `CallRuntime` explicito + - preparar o terreno para policies e scheduler unificados + +## Entrada 002 - Fase 2 / runtime explicito + +- Data: 2026-03-27 +- Objetivo: encapsular o fluxo da chamada em um `CallRuntime` com `CallState` explicito, reduzindo estado implícito e `nonlocal` em `app/livekit/main.py` +- Arquivos alterados: + - `app/livekit/main.py` + - `app/livekit/runtime/state.py` + - `app/livekit/runtime/call_runtime.py` + - `docs/refactor-log.md` + - `docs/api-overview.md` +- Comportamento preservado: + - contrato do websocket do bridge + - fluxo STT -> pipeline -> TTS + - idle nudge, interrupcao e finalizacao continuam com a mesma logica operacional + - exportacao e notificacao `DONE` continuam sendo disparadas pelo agent +- Comportamento alterado: + - `app/livekit/main.py` passa a ser majoritariamente wiring do session e construcao do runtime + - o estado mutavel da chamada fica concentrado em `CallState` +- Risco conhecido: + - as regras ainda nao foram transformadas em policies puras; a logica segue acoplada ao runtime, apenas mais organizada + - ainda faltam testes automatizados para cenarios concorrentes de fala, interrupcao e fechamento +- Validacao executada: + - revisao manual do diff + - `PYTHONDONTWRITEBYTECODE=1 python3 -m py_compile app/livekit/main.py app/livekit/adapters/bridge_gateway.py app/livekit/adapters/export_service.py app/livekit/adapters/pipeline_adapter.py app/livekit/adapters/speech_service.py app/livekit/runtime/state.py app/livekit/runtime/call_runtime.py` +- Proximo passo: + - introduzir command/execution boundaries no runtime + - preparar a migracao de interrupcao, idle e finalize para policies + scheduler + +## Entrada 003 - Fase 3 / commands e executor + +- Data: 2026-03-29 +- Objetivo: separar efeitos colaterais do `CallRuntime` por meio de comandos tipados e um executor dedicado +- Arquivos alterados: + - `app/livekit/main.py` + - `app/livekit/runtime/commands.py` + - `app/livekit/runtime/command_executor.py` + - `app/livekit/runtime/call_runtime.py` + - `docs/refactor-log.md` + - `docs/api-overview.md` +- Comportamento preservado: + - contrato do websocket do bridge + - fluxo STT -> pipeline -> TTS + - sequencia de idle nudge, interrupcao e finalizacao + - integracao com LiveKit, pipeline e export sem troca de contrato +- Comportamento alterado: + - `CallRuntime` deixa de chamar diretamente bridge/export/speech/pipeline para os principais side effects + - side effects passam a trafegar por comandos (`commands.py`) executados por `RuntimeCommandExecutor` +- Risco conhecido: + - ainda existe conhecimento de regras dentro de `CallRuntime`; a extracao atual isola efeitos, mas nao transforma as regras em policies puras + - `agent._ready`, `agent._run_lock` e coordenacao de fila ainda pertencem ao runtime +- Validacao executada: + - revisao manual do diff + - validacao sintatica com `python3 -c 'from pathlib import Path; paths = [...]; [compile(Path(p).read_text(), p, "exec") for p in paths]'` +- Proximo passo: + - extrair policies de interrupcao, idle e finalizacao + - introduzir scheduler unico para timers e reduzir logica condicional no runtime + +## Entrada 004 - Fase 4 / policies e scheduler + +- Data: 2026-03-29 +- Objetivo: mover regras de interrupcao, idle e finalize-once para policies dedicadas e centralizar timers em um scheduler unico +- Arquivos alterados: + - `app/livekit/main.py` + - `app/livekit/policies/interrupt_policy.py` + - `app/livekit/policies/idle_policy.py` + - `app/livekit/policies/finalization_policy.py` + - `app/livekit/runtime/scheduler.py` + - `app/livekit/runtime/state.py` + - `app/livekit/runtime/call_runtime.py` + - `docs/refactor-log.md` + - `docs/api-overview.md` +- Comportamento preservado: + - contrato do websocket do bridge + - fluxo STT -> pipeline -> TTS + - timers de idle nudge e idle close continuam ativos + - finalizacao por `DONE`, `room_empty`, `shutdown_callback` e `no_user_response` continua existindo +- Comportamento alterado: + - `InterruptPolicy` passa a decidir interrupcao por stage e filtro de backchannel + - `IdlePolicy` passa a concentrar regras de nudge/close + - `FinalizationPolicy` passa a concentrar regras de finalize-once e room-empty + - `TimerScheduler` passa a concentrar arm/cancel/token dos timers + - `CallState` deixa de carregar estado interno de timers +- Risco conhecido: + - o `CallRuntime` ainda conhece a ordem operacional completa da chamada; as policies reduzem acoplamento, mas ainda nao existe reducer/event engine + - `RuntimeConfig` ainda carrega campos hoje parcialmente sobrepostos pelas policies; isso pode ser simplificado em um corte futuro +- Validacao executada: + - revisao manual do diff + - validacao sintatica com `python3 -c 'from pathlib import Path; paths = [...]; [compile(Path(p).read_text(), p, "exec") for p in paths]'` +- Proximo passo: + - avaliar se vale introduzir eventos/comandos mais declarativos no runtime ou parar aqui e estabilizar + - se seguir, o proximo salto natural e um reducer/event engine ou a evolucao da integracao de IA (`llm_node` ou agente remoto) + +## Entrada 005 - Estabilizacao / testes unitarios do runtime + +- Data: 2026-03-29 +- Objetivo: adicionar cobertura automatizada para os cenarios mais criticos do runtime sem depender de LiveKit real +- Arquivos alterados: + - `tests/__init__.py` + - `tests/livekit/__init__.py` + - `tests/livekit/test_runtime.py` + - `docs/refactor-log.md` +- Comportamento preservado: + - nenhum contrato publico da API foi alterado + - nenhum fluxo de audio ou websocket foi modificado +- Comportamento alterado: + - o repositorio passa a ter uma suite unitária para policies, scheduler, command executor e cenarios críticos do `CallRuntime` +- Risco conhecido: + - os testes ainda usam doubles/fakes e nao substituem validacao integrada com LiveKit real + - ainda faltam cenarios mais completos de timer real, callback de room e speech handle real +- Validacao executada: + - `python3 -m unittest tests.livekit.test_runtime -v` + - `python3 -m unittest discover -s tests -v` +- Proximo passo: + - ampliar cobertura para speech handle real e integracao de callbacks com objetos LiveKit reais, se isso passar a ser area de regressao + - ou encerrar a fase de estabilizacao e voltar a discutir a evolucao da camada de IA + +## Entrada 006 - Backend remoto via websocket + +- Data: 2026-03-29 +- Objetivo: permitir trocar a pipeline local por um agent remoto via websocket, preservando o fluxo STT -> IA -> TTS e o runtime atual +- Arquivos alterados: + - `app/livekit/main.py` + - `app/livekit/adapters/agent_backend.py` + - `app/livekit/adapters/backend_factory.py` + - `app/livekit/adapters/pipeline_adapter.py` + - `app/livekit/adapters/remote_agent_ws_adapter.py` + - `app/livekit/runtime/command_executor.py` + - `requirements.txt` + - `tests/livekit/test_remote_agent_ws_adapter.py` + - `docs/refactor-log.md` + - `docs/api-overview.md` +- Comportamento preservado: + - contrato do websocket do bridge + - fluxo de audio com LiveKit e TTS no agent local + - runtime de interrupcao, idle nudge e finalizacao + - backend `langgraph` segue como default +- Comportamento alterado: + - o backend da camada de IA passa a ser selecionavel por `AGENT_BACKEND` + - quando `AGENT_BACKEND=remote_ws`, o agent envia cada turno transcrito para um websocket remoto e vocaliza a resposta retornada + - a finalizacao pode opcionalmente buscar um `result` remoto via mensagem `type=end` + - o bridge agora repassa `agent`, `RouterCallKeyDay`, `RouterCallKey`, `ANI`, `GSM` e `ID_FATURA` para o agent quando esses campos vierem no metadata do cliente + - o adapter remoto roteia a chamada para endpoints diferentes de acordo com `agent` (`conta`, `oferta`, `cobranca`) + - o agent `conta` passou a usar contrato proprio em websocket com `action/payload` na entrada e `result.content` na resposta +- Risco conhecido: + - o contrato do websocket remoto ainda e interno e precisa ser homologado com o agent externo real + - o adapter atual usa conexao websocket por turno, nao uma sessao persistente + - se o agent remoto nao devolver `stage`, a politica local dependera do ultimo stage conhecido ou do default configurado + - o campo `timestamp` e gerado localmente no momento de cada turno; se o integrador exigir outro formato, isso ainda precisa ser alinhado +- Validacao executada: + - testes unitarios do adapter remoto e factory + - execucao da suite `unittest` +- Proximo passo: + - homologar o contrato do websocket remoto com payload/resposta reais + - decidir se a conexao remota deve continuar por turno ou evoluir para sessao persistente + +## Entrada 007 - Cliente web de voz e call_config por chamada + +- Data: 2026-03-30 +- Objetivo: disponibilizar um cliente web para testar o fluxo de voz pelo navegador e permitir overrides de `STT`, `TTS` e `AGENT` por chamada +- Arquivos alterados: + - `app/ws_gateway/main.py` + - `app/ws_gateway/voice_client.html` + - `app/ws_gateway/call_config.py` + - `app/livekit/main.py` + - `app/livekit/call_config.py` + - `app/livekit/adapters/backend_factory.py` + - `tests/config/test_call_config.py` + - `docs/api-overview.md` + - `docs/refactor-log.md` +- Comportamento preservado: + - fluxo principal de audio via `/ws/agent` + - bridge continua recebendo `start`, audio binario e encerrando com `stop` + - defaults de `STT`, `TTS` e backend continuam vindo de env quando nao houver override +- Comportamento alterado: + - `/voice-client` agora serve um cliente web para capturar microfone e ouvir o audio retornado + - o `start` pode carregar `call_config` com overrides por chamada + - o agent agora aplica overrides de backend, STT e TTS para a sessao corrente +- Risco conhecido: + - o cliente web usa `ScriptProcessorNode`, que e suficiente para homologacao mas nao e a opcao mais moderna da Web Audio API + - os providers dinamicos suportados ainda sao os que ja existem no projeto (`internal_http` e `elevenlabs`) +- Validacao executada: + - testes puros de `call_config` + - execucao da suite `unittest` +- Proximo passo: + - se o cliente web passar a ser usado em rotina de homologacao, vale migrar a captura/playback para `AudioWorklet` + - se surgirem novos providers de STT/TTS, plugar nas factories por chamada + +## Entrada 008 - Timeline estruturada por chamada + +- Data: 2026-03-30 +- Objetivo: criar observabilidade ponta a ponta para entender o caminho `bridge -> livekit -> stt -> backend -> tts -> stop` por chamada +- Arquivos alterados: + - `app/utils/call_timeline.py` + - `app/ws_gateway/main.py` + - `app/livekit/main.py` + - `app/livekit/runtime/call_runtime.py` + - `app/livekit/adapters/backend_factory.py` + - `app/livekit/adapters/pipeline_adapter.py` + - `app/livekit/adapters/remote_agent_ws_adapter.py` + - `app/livekit/adapters/bridge_gateway.py` + - `app/providers/stt_internal_livekit.py` + - `tests/utils/test_call_timeline.py` + - `docs/api-overview.md` + - `docs/refactor-log.md` +- Comportamento preservado: + - contrato publico de `WS /ws/agent` + - fluxo de audio ja existente entre cliente, bridge, LiveKit e agent + - backends de IA, STT e TTS seguem com a mesma responsabilidade funcional +- Comportamento alterado: + - bridge e agent agora escrevem uma timeline estruturada em `jsonl` compartilhada por chamada + - a timeline usa o mesmo `timeline_id` e o mesmo `origin_unix_ms` entre os dois processos para manter a ordem relativa dos eventos + - o bridge passou a registrar eventos como `ready_sent`, `dispatch_started`, `agent_join`, `client_audio_enabled`, `done_packet_received` e `call_end` + - o agent passou a registrar eventos como `user_transcript_final`, `pipeline_run_started`, `pipeline_run_completed`, `tts_stage_started`, `interrupt_marked`, `finalize_started` e `finalize_completed` + - o provider de STT agora registra etapas de reconhecimento (`stt_recognize_started`, `stt_vosk_completed`, `stt_http_completed`) + - o backend remoto por websocket agora registra `remote_agent_request` e `remote_agent_response` +- Risco conhecido: + - a timeline registra payloads de request do backend remoto, o que aumenta a verbosidade e pode expor dados sensiveis em ambiente de debug + - o arquivo `jsonl` e compartilhado entre dois processos locais; o uso atual com `flock` e suficiente para dev/homologacao, mas nao substitui observabilidade centralizada +- Validacao executada: + - `py_compile` dos arquivos alterados + - teste unitario novo de timeline + - execucao completa da suite `unittest` +- Proximo passo: + - decidir se a timeline deve continuar sempre habilitada em dev ou ficar atras de flag por ambiente + - se a homologacao exigir, adicionar visualizador simples da timeline no cliente web + +## Entrada 009 - Fake remoto para homologacao via UI + +- Data: 2026-03-30 +- Objetivo: permitir testar o fluxo `STT -> remote_ws -> TTS` pela UI mesmo sem o agent remoto real estar de pe +- Arquivos alterados: + - `app/ws_gateway/fake_remote_agent.py` + - `app/ws_gateway/main.py` + - `app/ws_gateway/voice_client.html` + - `app/livekit/adapters/backend_factory.py` + - `tests/ws_gateway/test_fake_remote_agent.py` + - `tests/livekit/test_remote_agent_ws_adapter.py` + - `docs/api-overview.md` + - `docs/refactor-log.md` +- Comportamento preservado: + - backend `remote_ws` continua exigindo URL real quando selecionado + - o contrato do adapter remoto segue o mesmo +- Comportamento alterado: + - o `bridge` agora expoe `WS /fake-agent/ws` + - a UI ganhou a opcao `remote_ws_fake` + - quando `remote_ws_fake` e selecionado, o agent local aponta para `REMOTE_AGENT_WS_FAKE_URL` ou usa por padrao `ws://127.0.0.1:8000/fake-agent/ws` + - o fake responde nos contratos `conta` e generico, com progressao simples de stages e encerramento por palavras-chave como `encerrar`, `obrigado` e `tchau` +- Risco conhecido: + - o fake nao simula comportamento de negocio real; ele serve apenas para homologacao tecnica do fluxo de voz + - o fake e stateless por conexao, entao a progressao depende do `stage` enviado pelo agent local +- Validacao executada: + - testes unitarios do fake remoto + - testes do backend fake no adapter remoto + - execucao completa da suite `unittest` +- Proximo passo: + - se a homologacao pedir cenarios mais realistas, adicionar scripts por agent (`conta`, `oferta`, `cobranca`) com respostas configuraveis + +## Template de nova entrada + +- Data: +- Objetivo: +- Arquivos alterados: +- Comportamento preservado: +- Comportamento alterado: +- Risco conhecido: +- Validacao executada: +- Proximo passo: + +## Entrada 010 - Correcao do roteamento `remote_ws_fake` + +- Data: 2026-03-30 +- Objetivo: garantir que a selecao `remote_ws_fake` na UI nao seja sobrescrita pela URL real configurada em `REMOTE_AGENT_WS_URL` +- Arquivos alterados: + - `app/livekit/adapters/backend_factory.py` + - `app/livekit/adapters/remote_agent_ws_adapter.py` + - `tests/livekit/test_remote_agent_ws_adapter.py` + - `docs/api-overview.md` + - `docs/refactor-log.md` +- Comportamento preservado: + - `remote_ws` continua usando `REMOTE_AGENT_WS_URL` + - o adapter remoto continua sendo o mesmo para backend real e fake +- Comportamento alterado: + - `remote_ws_fake` agora prioriza sempre `REMOTE_AGENT_WS_FAKE_URL`, mesmo quando `REMOTE_AGENT_WS_URL` existe no ambiente + - a timeline do adapter remoto agora diferencia `backend=remote_ws_fake` de `backend=remote_ws` +- Risco conhecido: + - aliases como `fake_remote_ws` e `ws_fake` continuam sendo tratados como `remote_ws_fake`, mas o label emitido na timeline fica canonico como `remote_ws_fake` +- Validacao executada: + - `python3 -m unittest tests.livekit.test_remote_agent_ws_adapter -v` + - `python3 -m py_compile app/livekit/adapters/backend_factory.py app/livekit/adapters/remote_agent_ws_adapter.py tests/livekit/test_remote_agent_ws_adapter.py` +- Proximo passo: + - validar na chamada real da UI que o timeline e o request usam `ws://127.0.0.1:8000/fake-agent/ws` + +## Entrada 011 - Reorganizacao da raiz e remocao de boilerplate + +- Data: 2026-04-06 +- Objetivo: consolidar a estrutura real do projeto na raiz do repositorio, mantendo `src/app` e `src/agent` como codigo-fonte canonico +- Arquivos alterados: + - `README.md` + - `makefile` + - `Dockerfile` + - `pytest.ini` + - `.env.example` + - `.gitignore` + - `.dockerignore` + - `requirements.txt` + - `livekit.yaml` + - `docs/*` + - `tests/*` + - `k8s/deployment.yaml` + - remocao de `pyproject.toml`, `uv.lock`, `README copy.md`, `configs/config.example.yaml` e scripts herdados do boilerplate +- Comportamento preservado: + - `src/app` e `src/agent` permanecem como raiz do codigo de execucao + - o bridge FastAPI continua em `src/app/ws_gateway/main.py` + - o projeto continua usando `PYTHONPATH=src` +- Comportamento alterado: + - a raiz do repositorio passa a refletir o projeto real, e nao mais o boilerplate original + - `requirements.txt` passa a ser a fonte principal de dependencias + - artefatos de runtime (`logs`, `timeline`, `__pycache__`) deixam de ser versionados +- Risco conhecido: + - arquivos de infra legados fora do fluxo principal ainda poderiam carregar suposicoes antigas do boilerplate + - a migracao para `requirements.txt` simplifica a operacao, mas remove o lockfile anterior +- Validacao executada: + - `python3 -m compileall src tests` + - `python3 -m pytest tests/utils/test_call_timeline.py -q` +- Proximo passo: + - revisar a infra auxiliar restante e reduzir ruido operacional na raiz + +## Entrada 012 - Simplificacao do compose local de observabilidade + +- Data: 2026-04-06 +- Objetivo: transformar o `docker-compose.yml` em uma stack minima e coerente para Langfuse local +- Arquivos alterados: + - `docker-compose.yml` +- Comportamento preservado: + - a stack opcional de Langfuse continua disponivel para uso local + - a aplicacao principal segue fora do compose, operada via `make` +- Comportamento alterado: + - remocao do servico `mongo` + - consolidacao apenas de `langfuse-web`, `langfuse-worker`, `postgres`, `redis`, `clickhouse` e `minio` + - reducao de duplicacao com anchors para configuracoes compartilhadas +- Risco conhecido: + - o compose continua sendo opcional e depende de configuracao explicita do projeto para apontar para Langfuse self-hosted +- Validacao executada: + - validacao sintatica do YAML com parser local +- Proximo passo: + - se a equipe mantiver uso frequente, considerar alvos dedicados no `makefile` para subir e derrubar a stack + +## Entrada 013 - Provider TTS opcional desacoplado do import global + +- Data: 2026-04-06 +- Objetivo: impedir que a falta do SDK do ElevenLabs quebrasse o import do modulo inteiro de TTS +- Arquivos alterados: + - `src/app/providers/tts.py` + - `tests/providers/test_tts.py` +- Comportamento preservado: + - selecao de provider TTS por configuracao + - suporte aos providers ja existentes +- Comportamento alterado: + - o SDK do ElevenLabs deixa de ser carregado no topo do modulo e passa a ser resolvido sob demanda + - ausencia do SDK passa a gerar `missing_elevenlabs_sdk`, em vez de falha de import que derruba testes e providers nao relacionados +- Risco conhecido: + - o provider continua opcional, mas a disciplina de dependencias ainda segue concentrada em `requirements.txt` +- Validacao executada: + - `python3 -m pytest tests/providers/test_tts.py -q` + - `python3 -m pytest tests/adapters/test_azure_tts.py -q` +- Proximo passo: + - continuar removendo acoplamentos desnecessarios entre providers e entrypoints + +## Entrada 014 - Remocao de mocks e hardcodes dos pipelines de texto + +- Data: 2026-04-06 +- Objetivo: tirar codigo de demonstracao do caminho produtivo dos endpoints de texto +- Arquivos alterados: + - `src/app/services/session_context.py` + - `src/app/services/text_pipeline.py` + - `src/app/services/text_pipeline_stream.py` + - `src/app/ws_gateway/main.py` + - `tests/services/test_session_context.py` +- Comportamento preservado: + - os endpoints `/ws/text` e `/ws/text_stream` continuam aceitando a mesma sessao de entrada + - os pipelines continuam recebendo `mailing` e `intro` +- Comportamento alterado: + - remocao do uso de `app.models.mock` + - remocao de protocolo e dados fixos hardcoded nos pipelines + - introducao de extracao explicita de protocolo via `session_context` +- Risco conhecido: + - se existiam fluxos de dev que dependiam dos mocks antigos, eles passam a precisar de payloads reais +- Validacao executada: + - `python3 -m pytest -q` + - resultado observado na etapa: `36 passed` +- Proximo passo: + - consolidar contratos duplicados entre gateway e runtime + +## Entrada 015 - Unificacao de `call_config` + +- Data: 2026-04-06 +- Objetivo: eliminar duplicacao entre as implementacoes de `call_config` do gateway e do runtime LiveKit +- Arquivos alterados: + - `src/app/common/call_config.py` + - `src/app/ws_gateway/call_config.py` + - `src/app/livekit/call_config.py` + - `tests/config/test_call_config.py` +- Comportamento preservado: + - API publica de `build_call_config(...)` + - API publica de `normalize_call_config(...)` e resolucao dos campos auxiliares +- Comportamento alterado: + - a normalizacao passa a ter um nucleo compartilhado em `src/app/common/call_config.py` + - `ws_gateway` e `livekit` viram wrappers leves para o mesmo contrato +- Risco conhecido: + - qualquer evolucao futura de `call_config` passa a ter impacto compartilhado entre bridge e agent, o que e desejado, mas exige disciplina de compatibilidade +- Validacao executada: + - `python3 -m pytest tests/config/test_call_config.py -q` + - `python3 -m pytest -q` + - resultado observado na etapa: `36 passed` +- Proximo passo: + - iniciar a quebra incremental do monolito em `src/app/ws_gateway/main.py` + +## Entrada 016 - Extracao do parsing inicial da sessao no bridge + +- Data: 2026-04-06 +- Objetivo: remover do `ws_gateway/main.py` a logica de parsing da mensagem `start`, resolucao de mailing e montagem de `intro`/`nudge` +- Arquivos alterados: + - `src/app/ws_gateway/session_start.py` + - `src/app/ws_gateway/main.py` + - `tests/ws_gateway/test_session_start.py` +- Comportamento preservado: + - contrato de primeira mensagem `type=start` + - geracao de `mailing`, `intro` e `nudge` + - construcao do contexto do agente remoto a partir de metadata + mailing +- Comportamento alterado: + - criacao da dataclass `StartSessionContext` + - centralizacao de `parse_start_payload`, `recv_start_message` e `build_remote_agent_context` em modulo proprio + - `/ws/agent`, `/ws/text` e `/ws/text_stream` passam a consumir um contexto estruturado, e nao uma tuple solta +- Risco conhecido: + - o contrato de entrada continua dependente do payload do cliente; a extracao melhora testabilidade, mas nao redefine o protocolo +- Validacao executada: + - `python3 -m compileall src/app/ws_gateway src/app/services tests/ws_gateway` + - `python3 -m pytest tests/ws_gateway/test_session_start.py -q` + - `python3 -m pytest -q` + - resultado observado na etapa: `39 passed` +- Proximo passo: + - extrair o bootstrap da chamada para reduzir ainda mais a responsabilidade do endpoint + +## Entrada 017 - Extracao do bootstrap da chamada do bridge + +- Data: 2026-04-06 +- Objetivo: separar a preparacao da chamada do runtime do endpoint `/ws/agent` +- Arquivos alterados: + - `src/app/ws_gateway/session_bootstrap.py` + - `src/app/ws_gateway/main.py` + - `tests/ws_gateway/test_session_bootstrap.py` +- Comportamento preservado: + - geracao de token de acesso ao LiveKit + - montagem de `call_config`, `remote_agent_context`, `timeline` e `dispatch_metadata` + - fallbacks de protocolo e telefone +- Comportamento alterado: + - introducao da dataclass `BridgeSessionBootstrap` + - centralizacao de `room_name`, `identity`, `token`, `protocol`, `phone_number` e `dispatch_metadata` em um builder dedicado + - `main.py` passa a consumir um pacote pronto de bootstrap +- Risco conhecido: + - o bootstrap ainda depende de envs e factories externas; a extracao organiza responsabilidade, mas nao muda a politica de configuracao +- Validacao executada: + - `python3 -m compileall src/app/ws_gateway tests/ws_gateway` + - `python3 -m pytest tests/ws_gateway/test_session_bootstrap.py -q` + - `python3 -m pytest -q` + - resultado observado na etapa: `41 passed` +- Proximo passo: + - extrair o ciclo de vida da sessao LiveKit do endpoint + +## Entrada 018 - Extracao do lifecycle da sessao LiveKit no bridge + +- Data: 2026-04-06 +- Objetivo: remover do endpoint os handlers de participante, o processamento do pacote `DONE` e o watcher de encerramento +- Arquivos alterados: + - `src/app/ws_gateway/session_lifecycle.py` + - `src/app/ws_gateway/main.py` + - `tests/ws_gateway/test_session_lifecycle.py` +- Comportamento preservado: + - deteccao de entrada e saida do agente remoto + - processamento do pacote `agent.stage` com `DONE` + - envio de `stop` para o cliente ao final da chamada +- Comportamento alterado: + - criacao de `RoomLifecycleState` para concentrar `agent_participant` e `done_payload` + - extracao de `register_room_lifecycle_handlers(...)` + - extracao de `watch_call_done(...)` +- Risco conhecido: + - o fluxo ainda depende de coordenacao concorrente entre tasks do bridge; a extracao reduz tamanho do endpoint, mas nao muda o modelo de concorrencia +- Validacao executada: + - `python3 -m compileall src/app/ws_gateway tests/ws_gateway` + - `python3 -m pytest tests/ws_gateway/test_session_lifecycle.py -q` + - `python3 -m pytest -q` + - resultado observado na etapa: `43 passed` +- Proximo passo: + - extrair a camada de intro/sinalizacao do bridge + +## Entrada 019 - Extracao da intro TTS e da sinalizacao de audio no bridge + +- Data: 2026-04-07 +- Objetivo: isolar do endpoint o trecho responsavel por intro TTS, liberacao do audio do cliente, sinalizacao `client_audio_enabled` e bootstrap do streaming do agente +- Arquivos alterados: + - `src/app/ws_gateway/session_audio.py` + - `src/app/ws_gateway/main.py` + - `tests/ws_gateway/test_session_audio.py` +- Comportamento preservado: + - intro sintetizada antes da liberacao do audio do cliente + - emissao do controle `client_audio_enabled` para o agente + - inicializacao do worker que consome audio do agente quando ele fica pronto +- Comportamento alterado: + - extracao de `play_intro_to_client(...)` + - extracao de `notify_client_audio_enabled(...)` + - extracao de `stream_agent_audio_when_ready(...)` + - `main.py` passa a gerenciar explicitamente a task `t_control`, em vez de disparar a notificacao como fire-and-forget + - ajuste de typing para evitar import-time acidental de `audioop` em testes unitarios do modulo novo +- Risco conhecido: + - o caminho quente de audio continua centralizado em `main.py`; ainda faltam as primitivas de transporte para o arquivo deixar de ser monolitico +- Validacao executada: + - `python3 -m compileall src/app/ws_gateway tests/ws_gateway` + - `python3 -m pytest tests/ws_gateway/test_session_audio.py -q` + - `python3 -m pytest -q` + - resultado observado na etapa: `46 passed` +- Proximo passo: + - extrair as primitivas de transporte de audio e LiveKit (`ws_audio_receiver`, `publish_queue_to_livekit`, `connect_publish_livekit`, `stream_agent_audio_to_queue`) + +## Entrada 020 - Validacao do Dockerfile para CI/CD + +- Data: 2026-04-07 +- Objetivo: revisar o `Dockerfile` frente ao layout atual do projeto e corrigir o principal risco de compatibilidade para pipeline +- Arquivos alterados: + - `Dockerfile` +- Comportamento preservado: + - execucao do bridge via `uvicorn app.ws_gateway.main:app` + - exposicao da porta `8000` + - `PYTHONPATH=/app/src` + - healthcheck em `/health` +- Comportamento alterado: + - troca da base de `python:3.13-slim` para `python:3.12-slim` + - o ajuste foi necessario porque o projeto ainda depende de `audioop` em `src/app/ws_gateway/main.py` e `src/app/utils/background.py` +- Risco conhecido: + - o build real da imagem nao conseguiu ser concluido nesta maquina porque nem o daemon Docker nem a conexao efetiva do Podman estavam operacionais no momento da validacao + - portanto, a validacao foi estrutural e de coerencia do arquivo, nao um smoke test completo de container +- Validacao executada: + - revisao manual de `Dockerfile`, `.dockerignore`, `requirements.txt`, `makefile` e `k8s/deployment.yaml` + - tentativa de `docker build`, bloqueada por daemon indisponivel + - tentativa de `podman build`, bloqueada por conexao local com a VM do Podman +- Proximo passo: + - executar um build real da imagem assim que houver runtime de containers disponivel no host ou direto no pipeline diff --git a/docs/refactor-plan.md b/docs/refactor-plan.md new file mode 100644 index 0000000..a4337aa --- /dev/null +++ b/docs/refactor-plan.md @@ -0,0 +1,116 @@ +# Refactor Plan + +## Objetivo + +Fazer um refactor incremental da camada de runtime de voz sem quebrar: +- contrato do bridge websocket +- integracao com LiveKit +- pipeline atual de negocio +- comportamento de interrupcao, idle nudge e finalizacao + +## Problema atual + +Hoje o arquivo `app/livekit/main.py` concentra responsabilidades demais: +- wiring do AgentSession +- timers +- politica de interrupcao +- idle nudge +- finalizacao +- chamada da pipeline +- speak / wait_for_playout +- notificacao ao bridge + +Isso dificulta manutencao, teste e evolucao para alternativas futuras como: +- `llm_node` +- agente remoto via API +- politica de chamada mais previsivel + +## Direcao arquitetural + +Separar o runtime em camadas: +- adapters: wrappers dos componentes atuais (pipeline, speech, bridge, export) +- runtime: coordenacao da chamada +- domain: estado, eventos e comandos +- policies: regras puras de interrupcao, idle e finalizacao +- engine: reducer e scheduler + +## Principios + +- preservar comportamento antes de trocar mecanismo +- extrair boundaries antes de mudar a orquestracao +- isolar side effects +- tornar finalizacao idempotente +- introduzir estado explicito da chamada + +## Fases sugeridas + +### Fase 1: Boundaries + +Extrair sem mudar comportamento: +- `PipelineAdapter` +- `SpeechService` +- `BridgeGateway` +- `ExportService` + +Saida esperada: +- `app/livekit/main.py` menor +- regras ainda iguais as de hoje + +### Fase 2: Runtime explicito + +Introduzir: +- `CallState` +- `CallRuntime` +- comandos e eventos + +Saida esperada: +- menor uso de `nonlocal` +- caminhos de execucao mais faceis de seguir + +### Fase 3: Policies e scheduler + +Migrar: +- interrupcao +- idle nudge +- idle close +- finalize once + +Saida esperada: +- regras de chamada isoladas e testaveis + +### Fase 4: Evolucao de IA + +Avaliar depois da estabilizacao: +- `llm_node` +- agente remoto via API +- SSE / websocket para agente externo + +Saida esperada: +- STT / IA / TTS com fronteiras mais limpas + +## Fora de escopo inicial + +- reescrever bridge +- substituir pipeline de negocio +- trocar protocolo websocket externo +- mudar contrato de audio + +## Definition of done por fase + +### Fase 1 +- sem mudanca de contrato externo +- sem mudanca de fluxo de audio +- classes novas usadas pelo `main.py` + +### Fase 2 +- estado da chamada centralizado +- menos logica de coordenacao inline + +### Fase 3 +- timers centralizados +- finalizacao unica e previsivel +- regras de interrupcao sem duplicacao + +### Fase 4 +- nova integracao de IA desacoplada do fluxo manual atual +- comparacao controlada com comportamento anterior diff --git a/docs/regional/ARQUITETURA_TIA_XAI_REGIONAL.md b/docs/regional/ARQUITETURA_TIA_XAI_REGIONAL.md new file mode 100644 index 0000000..2744235 --- /dev/null +++ b/docs/regional/ARQUITETURA_TIA_XAI_REGIONAL.md @@ -0,0 +1,416 @@ +# Arquitetura TIA Regional com Pool xAI Pré-Aquecido + +## 1. Introdução + +Este documento descreve a evolução do TIA para operar TTS xAI em alta disponibilidade e alta volumetria usando réplicas Kubernetes distribuídas por região. A solução parte da implementação atual do TIA/LiveKit e mantém seu contrato de TTS, suas métricas de underflow/TTFB e sua lógica de proteção contra repetição após áudio parcial. + +A mudança principal é mover a manutenção do conjunto de conexões WebSocket xAI para um **pool persistente por Pod TIA**, implementado como sidecar. Cada Pod regional pode manter até `XAI_POOL_SIZE` conexões xAI pré-aquecidas e prontas para uso. O TIA continua enxergando um endpoint compatível com xAI, porém local (`127.0.0.1`). + +## 2. Dificuldade do modelo anterior + +Na versão anterior, cada instância `OraclexAITTS` mantinha uma única conexão reutilizável. Como o LiveKit executa chamadas em processos de job, essas conexões são naturalmente distribuídas por chamadas/processos e não formam um pool global do Pod. + +Isso cria alguns riscos em alta volumetria: + +1. **Burst de handshakes WebSocket.** Muitas chamadas podem abrir conexões ao xAI praticamente ao mesmo tempo. +2. **Capacidade não compartilhada.** Uma chamada pode manter uma conexão ociosa enquanto outra precisa abrir uma nova. +3. **Ausência de backpressure no balanceador.** O LB enxerga o TIA como saudável mesmo quando a capacidade de TTS daquele Pod está totalmente ocupada. +4. **Acoplamento entre sessão e conexão xAI.** A quantidade de chamadas pode virar, indiretamente, a quantidade de conexões abertas, mesmo quando apenas uma parte das chamadas está sintetizando naquele instante. +5. **Falha regional afeta novas chamadas.** Sem readiness orientada à saúde/capacidade do xAI, novas sessões podem continuar chegando a uma réplica cuja região está degradada. + +O histórico de testes do TTS já mostrou que problemas de estabelecimento de WebSocket e jitter podem se manifestar de forma regional e em função de carga. A arquitetura proposta transforma esses sinais em capacidade operacional do Pod. + +## 3. Objetivos + +A solução tem os seguintes objetivos: + +- manter conexões OCI xAI abertas e pré-aquecidas; +- reutilizar uma conexão entre diferentes sínteses; +- reservar conexão upstream apenas durante uma utterance; +- liberar a conexão imediatamente após `audio.done`; +- impedir bursts de abertura de WebSockets no caminho crítico da chamada; +- permitir TIA ativo/ativo em múltiplas regiões; +- retirar automaticamente uma réplica saturada do balanceamento de novas conexões; +- manter chamadas existentes durante draining/rollout; +- renovar sockets antes do TTL de forma escalonada, evitando reconexão simultânea; +- preservar o protocolo atual do `OraclexAITTS` e minimizar mudanças no código de voz; +- permitir escalabilidade horizontal em Kubernetes. + +## 4. Arquitetura proposta + +```text + Clientes / Telefonia + | + v + Load Balancer / Service + (novas conexões WebSocket) + | + +--------------------+--------------------+ + | | + v v + Deployment ORD Deployment IAD + | | + +--------+--------+ +--------+--------+ + | Pod TIA ORD | | Pod TIA IAD | + | | | | + | bridge | | bridge | + | agent/livekit | | agent/livekit | + | | | | | | + | v | | v | + | xAI pool sidecar| | xAI pool sidecar| + | 50 WS warm | | 50 WS warm | + +-------+---------+ +-------+---------+ + | | + v v + OCI xAI ORD OCI xAI IAD +``` + +O mesmo `Service` Kubernetes seleciona Pods ORD e IAD. Cada deployment acrescenta o label `tia-region`, útil para métricas e operação, mas ambos compartilham `app=${APP_NAME}-regional`. + +## 5. Componentes + +### 5.1 Bridge TIA + +Mantém o WebSocket de entrada e o comportamento existente do TIA. O código principal não precisa conhecer o endpoint xAI regional. + +### 5.2 LiveKit Agent + +Continua instanciando `OraclexAITTS`, porém passa a usar: + +```text +XAI_WEBSOCKET_URL=ws://127.0.0.1:18100/xai/v1/tts +``` + +O adapter continua enviando: + +```text +text.clear +text.delta +text.done +``` + +e recebendo: + +```text +audio.clear +audio.delta +audio.done +``` + +Portanto, a lógica atual de streaming, TTFB, gaps, underflow e prevenção de replay permanece válida. + +### 5.3 Sidecar `xai-pool` + +Implementado em: + +```text +src/app/livekit/adapters/xai_pool_proxy.py +``` + +Responsabilidades: + +- abrir `XAI_POOL_SIZE` WebSockets no startup; +- usar abertura em ondas controladas (`XAI_POOL_PREWARM_CONCURRENCY`); +- manter os sockets vivos; +- recuperar automaticamente slots desconectados; +- renovar sockets antes do TTL; +- aplicar jitter no refresh para não reconectar todos simultaneamente; +- emprestar uma conexão a uma utterance; +- devolver a conexão ao pool depois de `audio.done`; +- expor health, readiness, status, drain e métricas. + +### 5.4 OCI xAI regional + +Cada deployment recebe seu endpoint próprio por `XAI_POOL_UPSTREAM_URL`. + +Exemplo: + +```text +ORD -> wss://...us-chicago-1.../xai/v1/tts +IAD -> wss://...us-ashburn-1.../xai/v1/tts +``` + +As credenciais reais ficam somente no sidecar de pool. O container `agent` usa uma credencial local dummy porque se conecta apenas a `localhost`. + +### 5.5 Kubernetes Service / Load Balancer + +O Service seleciona todas as réplicas regionais. O WebSocket funciona normalmente através do LB: o balanceador escolhe um Pod durante o HTTP Upgrade e aquela conexão permanece no mesmo backend durante sua vida. + +A estratégia de capacidade não tenta migrar uma conexão existente. Ela afeta **novas conexões**. + +## 6. Pool compartilhado por Pod + +Uma chamada não reserva um socket xAI por toda sua duração. + +```text +Call A falando ---------- sem TTS +Call B ouvindo ---------- sem nova síntese +Call C precisa falar ---- acquire WS #17 + text.clear + text.delta + text.done + audio.delta... + audio.done + release WS #17 +``` + +Portanto, 200 chamadas podem coexistir em um Pod com pool de 50, desde que não existam mais de 50 sínteses concorrentes naquele instante. + +Essa é a principal diferença entre dimensionar por **chamadas simultâneas** e por **utterances TTS simultâneas**. + +## 7. Readiness orientada à capacidade + +O sidecar expõe: + +```text +GET /healthz +GET /readyz +GET /pool/status +GET /metrics +POST /drain +``` + +`/healthz` responde se o processo está vivo. + +`/readyz` responde se o Pod deve aceitar **novas chamadas**. + +Exemplo padrão: + +```text +XAI_POOL_SIZE=50 +XAI_POOL_UNAVAILABLE_FREE=2 +XAI_POOL_RECOVER_FREE=5 +``` + +Com o Pod inicialmente pronto: + +```text +free > 2 -> ready +free <= 2 -> not ready (HTTP 503) +``` + +Depois de sair da rotação, só retorna quando: + +```text +free >= 5 -> ready novamente +``` + +Isso cria histerese e evita flapping de readiness em torno do limite. + +Para a política estrita sugerida de somente sair quando todas estiverem ocupadas: + +```text +XAI_POOL_UNAVAILABLE_FREE=0 +XAI_POOL_RECOVER_FREE=5 +``` + +Em produção recomenda-se uma pequena reserva (por exemplo 2 a 5 conexões), pois ela absorve rajadas e reduz a chance de uma sessão recém-chegada não encontrar capacidade. + +## 8. Como o LB redireciona tráfego + +Em Kubernetes, um Pod é Ready apenas quando todos os containers que possuem readiness probe estão Ready. + +O sidecar `xai-pool` possui uma probe em `/readyz`. + +Quando a capacidade acaba: + +```text +xai-pool /readyz -> 503 + | + v +Pod becomes NotReady + | + v +Pod removed from Service Endpoints + | + v +LB stops sending NEW connections +``` + +As conexões WebSocket já estabelecidas não são redirecionadas e continuam no Pod enquanto o processo continuar disponível. + +## 9. Renovação escalonada das conexões + +Manter 50 sockets abertos indefinidamente sem renovação é arriscado porque serviços upstream normalmente aplicam TTL e renovação de autorização. + +A configuração padrão usa: + +```text +XAI_POOL_CONNECTION_TTL_S=540 +XAI_POOL_REFRESH_JITTER_S=45 +``` + +Cada conexão recebe um refresh deadline diferente: + +```text +WS01 -> ~501s +WS02 -> ~527s +WS03 -> ~509s +... +``` + +Somente conexões livres são renovadas. Isso evita um evento no qual 50 conexões expiram e fazem handshake simultaneamente. + +## 10. Alta disponibilidade regional + +Operação normal: + +```text +LB +|-- ORD Pod 1 -> ready +|-- ORD Pod 2 -> ready +|-- IAD Pod 1 -> ready +`-- IAD Pod 2 -> ready +``` + +Se ORD perder saúde/capacidade xAI, os slots começam a falhar e a quantidade de conexões saudáveis/livres cai. Quando a readiness cruza o threshold, os Pods ORD saem dos endpoints e novas chamadas passam a ser atendidas pelos Pods IAD disponíveis. + +Não há necessidade de alterar o cliente ou o LiveKit para escolher a região. + +## 11. Escala horizontal + +A capacidade teórica de pool é: + +```text +capacidade regional de sockets = replicas_region * XAI_POOL_SIZE +``` + +Exemplo: + +```text +ORD: 2 pods x 25 sockets = 50 +IAD: 2 pods x 25 sockets = 50 +TOTAL = 100 sockets prewarmed +``` + +ou, se a OCI conceder capacidade independente suficiente: + +```text +ORD: 2 pods x 50 = 100 +IAD: 2 pods x 50 = 100 +TOTAL = 200 +``` + +### Restrição crítica + +`replicas * XAI_POOL_SIZE` **não pode ultrapassar o limite real concedido pela OCI para o endpoint/tenancy/região**. + +Se a OCI disser que o limite 50 é global por endpoint, então duas réplicas de 50 no mesmo endpoint seriam incorretas. Nesse caso use, por exemplo: + +```text +2 replicas x 25 = 50 total +``` + +ou obtenha endpoints/capacidades independentes. + +## 12. Escala visual + +```text +Carga baixa +========= +LB + |-- ORD-1 [pool 50: 10 leased / 40 free] + `-- IAD-1 [pool 50: 8 leased / 42 free] + +Carga aumenta +============= +LB + |-- ORD-1 [47 leased / 3 free] READY + `-- IAD-1 [30 leased /20 free] READY + +ORD satura +=========== +LB + |-- ORD-1 [48 leased /2 free] NOT READY -> sem novas chamadas + `-- IAD-1 [31 leased /19 free] READY -> recebe novas chamadas + +ORD recupera +============= +ORD-1 chega a 45 leased /5 free +/readyz volta a 200 +LB volta a considerá-lo para novas conexões +``` + +## 13. Draining e rollout + +O deployment usa: + +```yaml +strategy: + type: RollingUpdate + rollingUpdate: + maxUnavailable: 0 + maxSurge: 1 +``` + +O sidecar executa no `preStop`: + +```text +POST /drain +``` + +Isso torna `/readyz` imediatamente 503, removendo o Pod da entrada de novas conexões antes de sua finalização. + +`terminationGracePeriodSeconds` deve ser compatível com a duração/grace desejada para as chamadas existentes. + +## 14. Segurança + +As credenciais OCI/xAI reais ficam no Secret indicado por `XAI_SECRET_NAME` e são montadas apenas no sidecar. + +O agent conecta a localhost e não precisa conhecer a chave real do xAI. + +Para produção recomenda-se evoluir para `OKE_WORKLOAD_IDENTITY` sempre que suportado pela política do ambiente, eliminando API keys estáticas. + +## 15. Observabilidade + +`GET /pool/status` retorna: + +```json +{ + "status": "ready", + "region": "ord", + "configured": 50, + "healthy": 50, + "leased": 17, + "free": 33, + "total_acquires": 845, + "total_acquire_timeouts": 0, + "total_proxy_failures": 0 +} +``` + +`GET /metrics` expõe métricas Prometheus simples: + +- `tia_xai_pool_connections{state="healthy"}` +- `tia_xai_pool_connections{state="leased"}` +- `tia_xai_pool_connections{state="free"}` +- `tia_xai_pool_acquires_total` +- `tia_xai_pool_acquire_timeouts_total` +- `tia_xai_pool_proxy_failures_total` + +Estas métricas devem ser correlacionadas com as métricas já existentes no TIA, especialmente TTFB, gap e underflow. + +## 16. O que esta versão resolve + +| Problema | Solução | +|---|---| +| handshakes xAI no caminho crítico | pool pré-aquecido | +| burst de WebSockets | prewarm em ondas + refresh com jitter | +| conexão presa a uma chamada | lease somente durante utterance | +| LB envia tráfego a Pod sem capacidade TTS | `/readyz` baseado no pool | +| flapping de health | histerese unavailable/recover | +| indisponibilidade regional | Deployments ORD/IAD no mesmo Service | +| rollout derruba novas sessões | drain + RollingUpdate | +| chave xAI em todos os processos | credencial real somente no sidecar | +| observabilidade de capacidade | `/pool/status` + `/metrics` | + +## 17. O que esta versão não resolve sozinha + +A solução não cria capacidade de inferência no OCI xAI. Se todos os endpoints regionais terminarem no mesmo pool de inferência saturado, o TIA terá failover e melhor utilização de sockets, mas não multiplicará a capacidade real do modelo. + +É necessário confirmar com a OCI: + +1. limite de WebSockets por endpoint/região/tenancy; +2. se endpoints dedicados possuem capacidade independente; +3. se API servers/LB e inference workers são dedicados ou compartilhados; +4. quais limites podem ser reservados/negociados para a TIM. diff --git a/docs/regional/CHANGELOG_IMPLEMENTACAO.md b/docs/regional/CHANGELOG_IMPLEMENTACAO.md new file mode 100644 index 0000000..079f06f --- /dev/null +++ b/docs/regional/CHANGELOG_IMPLEMENTACAO.md @@ -0,0 +1,32 @@ +# Changelog — TIA Regional / xAI Pool + +## Implementação adicionada + +- sidecar `xai_pool_proxy` compatível com o protocolo WebSocket xAI usado pelo TIA; +- pool pré-aquecido configurável por Pod (`XAI_POOL_SIZE`, default de exemplo 50); +- acquire por utterance após `text.clear`; +- release automático após `audio.done`; +- recuperação de slots desconectados; +- refresh escalonado por TTL + jitter; +- prewarm em ondas para evitar burst de handshake; +- `/healthz`, `/readyz`, `/pool/status`, `/metrics` e `/drain`; +- readiness com histerese; +- Deployments regionais ORD/IAD usando a mesma imagem; +- credencial xAI real isolada no sidecar; +- Service único para balancear Pods regionais; +- RollingUpdate, PDB e HPA; +- scripts de renderização, validação e deployment; +- manuais de arquitetura, implantação e testes. + +## Compatibilidade + +O modo anterior permanece disponível. O deployment regional é opcional e não remove `XAI_WEBSOCKET_URL` direto para OCI xAI. + +## Validações executadas na geração + +- `py_compile` / `compileall` do novo sidecar: PASS; +- parsing YAML dos templates Kubernetes: PASS; +- renderização ORD/IAD via `envsubst`: PASS; +- parsing YAML dos manifests renderizados: PASS. + +Não foi executado teste real contra OCI xAI, pois depende das credenciais/endpoints do ambiente TIM/OCI. O manual descreve smoke, saturação, failover e stress test a executar em FQA. diff --git a/docs/regional/DEPLOYMENT_TIA_XAI_REGIONAL.md b/docs/regional/DEPLOYMENT_TIA_XAI_REGIONAL.md new file mode 100644 index 0000000..1fb5edb --- /dev/null +++ b/docs/regional/DEPLOYMENT_TIA_XAI_REGIONAL.md @@ -0,0 +1,385 @@ +# Deployment — TIA Regional com Pool xAI + +## 1. Pré-requisitos + +- cluster Kubernetes/OKE funcional; +- `kubectl` configurado; +- `envsubst` (`gettext`) instalado na máquina de deployment; +- imagem TIA construída a partir desta versão; +- ConfigMap `${APP_NAME}-config` já existente; +- Secrets `${APP_NAME}-api-secrets`, `${APP_NAME}-google-sa-secret` e `shared-tls-secret` já existentes conforme deployment atual; +- endpoint e credencial xAI de cada região; +- quota xAI validada para o número total de WebSockets configurado. + +## 2. Arquivos adicionados + +```text +src/app/livekit/adapters/xai_pool_proxy.py + +k8s/regional/ + deployment-region.yaml + service.yaml + pdb.yaml + hpa-region.yaml + xai-secret.example.yaml + regions.env.example + +scripts/ + render-regional-k8s.sh + validate-regional-k8s.sh + deploy-regional-k8s.sh +``` + +## 3. Construir a imagem + +Use o pipeline existente ou Dockerfile atual do TIA. O novo sidecar utiliza a mesma imagem e apenas muda o módulo executado: + +```text +python -m app.livekit.adapters.xai_pool_proxy +``` + +Exemplo local: + +```bash +docker build -f k8s/tia/Dockerfile \ + -t iad.ocir.io//tia:regional-xai-pool-v1 . +``` + +Faça push para o registry usado pelo cluster. + +## 4. Criar arquivo de ambiente de deployment + +```bash +cp k8s/regional/regions.env.example k8s/regional/regions.env +``` + +Edite no mínimo: + +```bash +APP_NAME=tim-ai-atend-agnt-integ-tia +K8S_NAMESPACE=... +IMAGE_REPOSITORY=... +IMAGE_TAG=... + +ORD_XAI_UPSTREAM_URL=wss://...us-chicago-1.../xai/v1/tts +ORD_XAI_SECRET_NAME=xai-ord-credentials + +IAD_XAI_UPSTREAM_URL=wss://...us-ashburn-1.../xai/v1/tts +IAD_XAI_SECRET_NAME=xai-iad-credentials +``` + +## 5. Definir tamanho do pool + +Para um Pod com 50 sockets: + +```bash +XAI_POOL_SIZE=50 +``` + +Readiness recomendada: + +```bash +XAI_POOL_UNAVAILABLE_FREE=2 +XAI_POOL_RECOVER_FREE=5 +``` + +Se desejar exatamente o comportamento "só sair quando os 50 estiverem ocupados": + +```bash +XAI_POOL_UNAVAILABLE_FREE=0 +XAI_POOL_RECOVER_FREE=5 +``` + +### Importante + +Se houver `N` réplicas apontando para o mesmo endpoint: + +```text +sockets máximos = N * XAI_POOL_SIZE +``` + +Nunca configure isso acima da quota xAI real. + +## 6. Criar Secrets regionais + +Não versione chaves reais. + +Exemplo por linha de comando: + +```bash +kubectl -n "$K8S_NAMESPACE" create secret generic xai-ord-credentials \ + --from-literal=XAI_API_KEY='' \ + --from-literal=OCI_COMPARTMENT_ID='' + +kubectl -n "$K8S_NAMESPACE" create secret generic xai-iad-credentials \ + --from-literal=XAI_API_KEY='' \ + --from-literal=OCI_COMPARTMENT_ID='' +``` + +Para `API_KEY`, `OCI_COMPARTMENT_ID` pode ficar vazio se o fluxo upstream não o exigir. + +Para produção, prefira OCI Vault/External Secrets/Workload Identity em vez de gravar chaves no repositório. + +## 7. Renderizar manifests + +```bash +./scripts/render-regional-k8s.sh +``` + +Saída: + +```text +k8s/regional/rendered/ + deployment-ord.yaml + deployment-iad.yaml + hpa-ord.yaml + hpa-iad.yaml + service.yaml + pdb.yaml +``` + +É possível escolher outro arquivo e diretório: + +```bash +./scripts/render-regional-k8s.sh ./minha-config.env /tmp/tia-regional +``` + +## 8. Validar sem implantar + +```bash +./scripts/validate-regional-k8s.sh +``` + +O script usa: + +```bash +kubectl apply --dry-run=client +``` + +Revise também: + +```bash +kubectl diff -f k8s/regional/rendered/ +``` + +## 9. Implantar + +```bash +./scripts/deploy-regional-k8s.sh +``` + +Ou manualmente: + +```bash +kubectl apply -f k8s/regional/rendered/service.yaml +kubectl apply -f k8s/regional/rendered/pdb.yaml +kubectl apply -f k8s/regional/rendered/deployment-ord.yaml +kubectl apply -f k8s/regional/rendered/deployment-iad.yaml +kubectl apply -f k8s/regional/rendered/hpa-ord.yaml +kubectl apply -f k8s/regional/rendered/hpa-iad.yaml +``` + +## 10. Verificar rollout + +```bash +kubectl -n "$K8S_NAMESPACE" get pods -l app=${APP_NAME}-regional -o wide +``` + +Todos os containers devem ficar Ready: + +```text +bridge 1/1 +agent 1/1 +xai-pool 1/1 +``` + +Confira os deployments: + +```bash +kubectl -n "$K8S_NAMESPACE" rollout status deployment/${APP_NAME}-ord +kubectl -n "$K8S_NAMESPACE" rollout status deployment/${APP_NAME}-iad +``` + +## 11. Validar o pool + +Escolha um Pod: + +```bash +POD=$(kubectl -n "$K8S_NAMESPACE" get pod \ + -l app=${APP_NAME}-regional,tia-region=ord \ + -o jsonpath='{.items[0].metadata.name}') +``` + +Port-forward: + +```bash +kubectl -n "$K8S_NAMESPACE" port-forward "$POD" 18100:18100 +``` + +Em outro terminal: + +```bash +curl -s http://127.0.0.1:18100/healthz | jq +curl -s http://127.0.0.1:18100/readyz | jq +curl -s http://127.0.0.1:18100/pool/status | jq +curl -s http://127.0.0.1:18100/metrics +``` + +Resultado esperado após prewarm: + +```json +{ + "region": "ord", + "configured": 50, + "healthy": 50, + "leased": 0, + "free": 50, + "status": "ready" +} +``` + +## 12. Logs + +```bash +kubectl -n "$K8S_NAMESPACE" logs "$POD" -c xai-pool -f +``` + +Eventos relevantes: + +```text +XAI_POOL_SLOT_OPENED +XAI_POOL_PREWARM_FAILED +XAI_POOL_SLOT_RECOVERY_FAILED +XAI_POOL_STARTED +``` + +## 13. Testar saturação/readiness + +Objetivo: provar que o Pod sai da rotação para novas sessões quando o pool atinge o limite. + +1. Gere sínteses concorrentes suficientes para ocupar o pool. +2. Observe: + +```bash +watch -n 1 'curl -s http://127.0.0.1:18100/pool/status | jq' +``` + +3. Quando `free <= XAI_POOL_UNAVAILABLE_FREE`, espere: + +```bash +curl -i http://127.0.0.1:18100/readyz +``` + +Resultado: + +```text +HTTP/1.1 503 Service Unavailable +``` + +4. Verifique o Pod: + +```bash +kubectl get pod "$POD" +``` + +Ele deve ficar `NotReady` enquanto o sidecar estiver sem capacidade. + +5. Quando conexões forem liberadas e `free >= XAI_POOL_RECOVER_FREE`, o `/readyz` volta a 200 e o Pod retorna aos endpoints. + +## 14. Verificar endpoints do Service + +```bash +kubectl -n "$K8S_NAMESPACE" get endpoints ${APP_NAME}-regional -o wide +``` + +ou, em clusters novos: + +```bash +kubectl -n "$K8S_NAMESPACE" get endpointslice \ + -l kubernetes.io/service-name=${APP_NAME}-regional -o yaml +``` + +Um Pod NotReady não deve ser usado para novas conexões do Service. + +## 15. Testar failover regional + +### Teste controlado de ORD + +Coloque ORD em drain: + +```bash +ORD_POD=$(kubectl -n "$K8S_NAMESPACE" get pod \ + -l app=${APP_NAME}-regional,tia-region=ord \ + -o jsonpath='{.items[0].metadata.name}') + +kubectl -n "$K8S_NAMESPACE" exec "$ORD_POD" -c xai-pool -- \ + python -c "import urllib.request; urllib.request.urlopen(urllib.request.Request('http://127.0.0.1:18100/drain', method='POST')).read()" +``` + +Confirme `/readyz=503` e gere **nova chamada**. Ela deve ser entregue a um Pod IAD ainda Ready. + +Esse teste não deve ser interpretado como migração de uma chamada existente; somente novas conexões são rebalanceadas. + +## 16. Testar recuperação + +O drain manual é intencional e não é revertido. Para voltar o Pod, reinicie-o: + +```bash +kubectl -n "$K8S_NAMESPACE" delete pod "$ORD_POD" +``` + +O novo Pod deve: + +1. iniciar sidecar; +2. pré-aquecer sockets; +3. atingir `recover_free_threshold`; +4. ficar Ready; +5. entrar novamente nos endpoints. + +## 17. Load Balancer e WebSocket + +WebSocket é suportado porque a conexão começa como HTTP Upgrade. O LB escolhe o backend no handshake e mantém a conexão naquele Pod. + +Não use sticky session como requisito de correção. A conexão WebSocket em si já é persistente ao backend selecionado. Em caso de perda do Pod, o cliente precisa reconectar. + +Uma política equivalente a least-connections pode ser útil quando o LB permite configuração, mas o mecanismo primário de proteção desta arquitetura é a readiness baseada em capacidade real do TTS. + +## 18. HPA + +Os manifests incluem HPA por CPU como proteção inicial. + +Atenção: aumentar réplicas também aumenta o número de sockets xAI pré-aquecidos. Portanto, HPA só deve ter `maxReplicas` maior que 1 se a quota xAI permitir: + +```text +HPA_MAX_REPLICAS * XAI_POOL_SIZE <= quota regional permitida +``` + +Para evolução futura, prefira uma métrica customizada que considere chamadas/TTS ativos e também um controlador que respeite orçamento global de sockets. + +## 19. Rollback + +Para voltar ao modelo anterior: + +1. redirecione o Service/LB para o deployment antigo; +2. ou restaure o manifesto `k8s/tia/deployment.yaml` original; +3. o código antigo `OraclexAITTS` continua presente e compatível com endpoint xAI direto. + +A alteração não remove suporte ao modo anterior. + +## 20. Checklist de produção + +- [ ] confirmar quota real de WebSockets por endpoint/região/tenancy; +- [ ] confirmar se ORD e IAD têm pools de capacidade independentes; +- [ ] validar `XAI_POOL_SIZE * replicas` por região; +- [ ] validar Secret/Vault; +- [ ] medir tempo de prewarm dos 50 sockets; +- [ ] observar taxa de erro de handshake; +- [ ] validar refresh escalonado em janela > 10 minutos; +- [ ] testar saturation -> `/readyz=503`; +- [ ] testar recovery -> `/readyz=200`; +- [ ] testar drain durante chamada ativa; +- [ ] testar rollout sem perda de novas chamadas; +- [ ] testar perda total de ORD e entrada em IAD; +- [ ] correlacionar TTFB/gap/underflow com `pool_free`; +- [ ] validar LB/Service com WebSocket real; +- [ ] executar stress test com perfil semelhante ao tráfego de produção. diff --git a/docs/regional/TESTES_TIA_XAI_REGIONAL.md b/docs/regional/TESTES_TIA_XAI_REGIONAL.md new file mode 100644 index 0000000..677f55d --- /dev/null +++ b/docs/regional/TESTES_TIA_XAI_REGIONAL.md @@ -0,0 +1,151 @@ +# Plano de Testes — TIA Regional / Pool xAI + +## Objetivo + +Validar separadamente capacidade do pool, comportamento do Kubernetes, failover regional, impacto de latência e resiliência do xAI. + +## Camada 1 — Smoke + +1. subir 1 Pod ORD com pool pequeno (`XAI_POOL_SIZE=3`); +2. confirmar 3 conexões healthy; +3. realizar uma chamada e uma síntese; +4. confirmar que `leased` vai 0 -> 1 -> 0; +5. confirmar que o socket upstream continua healthy após `audio.done`. + +## Camada 2 — Prewarm + +Subir com `XAI_POOL_SIZE=50` e medir: + +- tempo total até Ready; +- taxa de sucesso de handshake; +- número máximo de handshakes simultâneos; +- impacto de `XAI_POOL_PREWARM_CONCURRENCY` 2, 5 e 10. + +Critério inicial: nenhuma rajada deve reproduzir os erros de conexão observados no modelo burst. + +## Camada 3 — Saturação + +Gerar mais utterances concorrentes que slots. + +Esperado: + +- leases nunca superam `XAI_POOL_SIZE`; +- acquire espera até `XAI_POOL_ACQUIRE_TIMEOUT_S`; +- readiness fica 503 quando free cruza threshold; +- novas chamadas deixam de entrar naquele Pod; +- chamadas existentes continuam. + +## Camada 4 — Histerese + +Com size=10: + +```text +UNAVAILABLE_FREE=2 +RECOVER_FREE=5 +``` + +Esperado: + +- free=2 -> NotReady; +- free=3/4 -> continua NotReady; +- free=5 -> Ready. + +## Camada 5 — Refresh + +Use TTL curto em FQA: + +```text +XAI_POOL_CONNECTION_TTL_S=60 +XAI_POOL_REFRESH_JITTER_S=20 +``` + +Observe por 5 minutos. + +Esperado: + +- sockets são renovados individualmente; +- não existe burst de 50 reconnects; +- slots ocupados não são renovados no meio da síntese; +- pool retorna ao tamanho configurado. + +## Camada 6 — Falha upstream + +Bloqueie ORD ou aponte temporariamente para endpoint inválido. + +Esperado: + +- slots ORD tornam-se unhealthy; +- maintenance tenta recuperação; +- readiness ORD cai; +- Service deixa de enviar novas chamadas a ORD; +- IAD continua Ready. + +## Camada 7 — Latência + +Compare três cenários: + +A. TIA -> xAI direto sem pool + +B. TIA -> localhost pool -> xAI com socket já aquecido + +C. TIA -> localhost pool -> xAI durante recuperação/abertura de socket + +Meça: + +- `provider_ttfb_ms`; +- `end_to_end_ttfb_ms`; +- `max_audio_delta_gap_ms`; +- `xai_underrun_estimado_ms`; +- connect time do upstream; +- `pool_free` e `pool_leased`. + +Hipótese: o hop localhost adiciona latência desprezível frente ao TTFB do provider, enquanto remove handshake xAI do caminho crítico na situação normal. + +## Camada 8 — Carga semelhante a produção + +Evite somente burst C=200. Use sockets persistentes e concorrência de síntese representativa da operação real, seguindo a metodologia que produziu resultados reprodutíveis nos testes anteriores. + +Rodar pelo menos: + +```text +20% da carga alvo +50% +80% +100% +120% por janela curta +``` + +## Camada 9 — Rollout + +Com chamadas ativas: + +```bash +kubectl rollout restart deployment/ +``` + +Validar: + +- Pod antigo entra em drain; +- não recebe novas conexões; +- Pod novo preaquece antes de ficar Ready; +- `maxUnavailable=0` preserva capacidade durante rollout. + +## Camada 10 — Failover regional + +1. ORD e IAD Ready; +2. iniciar tráfego contínuo; +3. tornar ORD NotReady; +4. verificar novas chamadas em IAD; +5. recuperar ORD; +6. verificar reentrada progressiva. + +## Evidências a guardar + +- logs do sidecar; +- `/pool/status` em intervalos de 1s; +- métricas Prometheus; +- logs TIA de TTFB/underflow; +- quantidade de endpoints Ready por região; +- distribuição de chamadas por Pod; +- erros de WebSocket upstream; +- timestamps de `audio.done`. diff --git a/docs/vad-pause-logs.md b/docs/vad-pause-logs.md new file mode 100644 index 0000000..2e712ea --- /dev/null +++ b/docs/vad-pause-logs.md @@ -0,0 +1,193 @@ +# VAD pause logs + +Este documento resume como analisar pausas de fala do usuario nos logs locais do agente LiveKit. + +## Onde procurar + +Os logs por chamada ficam em `logs/`. + +O arquivo costuma trazer os identificadores principais logo no inicio: + +```text +CALL_START | room=dev-room-75a205ec | protocol=PRT-20260409-0001 | session_id=hf05d7c2-1a1a-42e9-8651-ccd6351faff4 | bridge=ws-bridge-dev-2755fb53 +``` + +Use principalmente: + +- `room`: identifica a sala LiveKit e permite cruzar com `timeline/.jsonl`. +- `session_id`: identifica a conversa/sessao. +- `message_id`: identifica cada turno de usuario ou resposta do agente. +- `user_seq`: sequencia de falas finais do usuario. + +## Evento principal: `vad_user_pause` + +`vad_user_pause` indica que o wrapper de VAD registrou uma pausa relevante na fala do usuario. + +Exemplo: + +```text +FLOW | step=vad_user_pause | decision=end_of_speech | silence_ms=704 | pause_min_ms=700 | speech_duration_ms=896 | min_interrupt_ms=1600 | eligible=False | probability=0.000 | raw_speech_ms=0 | raw_silence_ms=0 +``` + +Campos essenciais: + +| Campo | Como usar | +| --- | --- | +| `decision` | Tipo de pausa detectada. `end_of_speech` e o principal para segmentacao real. `pause_threshold_reached` e um alerta intermediario. | +| `silence_ms` | Silencio observado pelo VAD, em milissegundos. E o campo mais importante para pausa. | +| `pause_min_ms` | Minimo de silencio necessario para fechar a fala. Vem de `LIVEKIT_VAD_MIN_SILENCE_DURATION_S`. | +| `speech_duration_ms` | Duracao do trecho de fala fechado pelo VAD. Trechos muito baixos indicam fala picotada. | +| `min_interrupt_ms` | Minimo de fala para considerar interrupcao do bot. Nao e o limite de pausa. | +| `eligible` | Se a fala passou de `min_interrupt_ms`. Mais util para barge-in do que para pausa. | +| `probability` | Probabilidade de fala no frame atual. Ajuda a entender ruido/atividade fraca. | +| `raw_speech_ms` | Acumulo bruto usado pelo VAD para detectar inicio de fala. | +| `raw_silence_ms` | Acumulo bruto usado pelo VAD para detectar silencio. | + +## Como interpretar pausas + +### Pausa que fechou fala + +Priorize `decision=end_of_speech`. + +Exemplo: + +```text +decision=end_of_speech | silence_ms=704 | pause_min_ms=700 | speech_duration_ms=896 +``` + +Leitura: + +- o VAD fechou o trecho apos observar `704ms` de silencio; +- o minimo configurado era `700ms`; +- como passou apenas `4ms` do minimo, a configuracao esta bem sensivel; +- se isso acontece varias vezes durante uma frase natural, a fala esta sendo segmentada cedo demais. + +### Alerta intermediario + +`decision=pause_threshold_reached` significa que o silencio passou do limite minimo antes do fechamento final. + +Exemplo: + +```text +decision=pause_threshold_reached | silence_ms=960 | pause_min_ms=700 | speech_duration_ms=0 | raw_speech_ms=32 +``` + +Esse evento deve ser lido com cuidado quando: + +- `speech_duration_ms=0`; +- `raw_speech_ms` e muito baixo, como `32`; +- `silence_ms` e muito alto, como dezenas de segundos. + +Nesses casos, pode ser artefato de estado acumulado ou ruido antes de uma fala real. Para confirmar impacto na conversa, cruze com `stt_final`. + +## Eventos auxiliares + +### `vad_speech_start` + +Indica inicio de fala detectado pelo VAD. + +```text +FLOW | step=vad_speech_start | speech_duration_ms=128 | min_interrupt_ms=1600 +``` + +Use para ver quando o sistema saiu de `listening` para `speaking`. + +### `vad_speech_end` + +Indica fechamento do trecho de fala. + +```text +FLOW | step=vad_speech_end | decision=too_short | speech_duration_ms=896 | eligible=False | silence_ms=704 | pause_min_ms=700 +``` + +Use junto com `vad_user_pause`. Se `speech_duration_ms` for baixo varias vezes, a fala pode estar sendo picotada. + +### `vad_interrupt_check` + +Indica que a fala atingiu duracao suficiente para interrupcao do bot. + +```text +FLOW | step=vad_interrupt_check | decision=eligible_by_duration | speech_duration_ms=1600 | min_interrupt_ms=1600 +``` + +Esse evento ajuda mais a analisar barge-in/interrupcao do bot do que pausa final de fala. + +### `stt_final` + +Mostra o texto final enviado como turno de usuario. + +```text +FLOW | step=stt_final | message_id=12345678-1234-4234-9234-123456789abc | user_seq=3 | text=Veio mais cara +``` + +Use este evento para validar o efeito real da segmentacao. Se uma frase natural virou varios `stt_final`, o VAD/STT segmentou demais. + +## Checklist de analise + +1. Encontre o `CALL_START` e anote `room`, `session_id` e `protocol`. +2. Filtre os eventos `vad_user_pause`. +3. Priorize `decision=end_of_speech`. +4. Compare `silence_ms` com `pause_min_ms`. +5. Verifique se `speech_duration_ms` esta muito baixo. +6. Cruze com os `stt_final` seguintes. +7. Se uma frase esperada virou varias mensagens curtas, a segmentacao esta agressiva. + +## Sinais de segmentacao agressiva + +Exemplo de fala esperada: + +```text +por que a minha fatura veio mais cara +``` + +Exemplo de saida segmentada: + +```text +stt_final | text=porque +stt_final | text=a minha fatura +stt_final | text=Veio mais cara +``` + +Se isso vier acompanhado de pausas assim: + +```text +vad_user_pause | decision=end_of_speech | silence_ms=704 | pause_min_ms=700 +``` + +provavelmente o limite de pausa esta fechando a fala cedo demais. + +## Parametro de ajuste + +O limite principal e: + +```env +LIVEKIT_VAD_MIN_SILENCE_DURATION_S=0.7 +``` + +Ele aparece no log como: + +```text +pause_min_ms=700 +``` + +Aumentar esse valor tende a juntar mais frases, porque o VAD espera mais silencio antes de encerrar a fala. O custo e aumentar a latencia percebida: o agente demora um pouco mais para responder depois que o usuario termina. + +Valores para teste manual: + +- `0.85`: ajuste conservador. +- `1.0`: tende a reduzir mais a segmentacao. +- `1.2`: pode ajudar em fala pausada, mas pode deixar a conversa lenta. + +## Regra pratica + +Para pausa, olhe primeiro: + +```text +decision + silence_ms + pause_min_ms +``` + +Para impacto na conversa, cruze com: + +```text +stt_final + user_seq + message_id +``` diff --git a/k8s/helm_deploy.yaml b/k8s/helm_deploy.yaml new file mode 100644 index 0000000..dc3d409 --- /dev/null +++ b/k8s/helm_deploy.yaml @@ -0,0 +1,4 @@ +apiVersion: v2 +name: ${APP_NAME} +description: helm para tia +version: ${IMAGE_TAG} \ No newline at end of file diff --git a/k8s/livekit/Dockerfile b/k8s/livekit/Dockerfile new file mode 100644 index 0000000..dd2b09e --- /dev/null +++ b/k8s/livekit/Dockerfile @@ -0,0 +1,7 @@ +FROM livekit/livekit-server:v1.11.0 + +COPY livekit.yaml /etc/livekit/livekit.yaml + +EXPOSE 7880 7881 7882 + +CMD ["--config", "/etc/livekit/livekit.yaml"] diff --git a/k8s/livekit/configmap.yaml b/k8s/livekit/configmap.yaml new file mode 100644 index 0000000..d3aff92 --- /dev/null +++ b/k8s/livekit/configmap.yaml @@ -0,0 +1,19 @@ +apiVersion: v1 +kind: ConfigMap +metadata: + name: ${APP_NAME}-livekit-config + namespace: ${K8S_NAMESPACE} +data: + livekit.yaml: | + port: 7880 + log_level: warn + rtc: + tcp_port: 7881 + udp_port: 7882 + port_range_start: 50000 + port_range_end: 60000 + redis: + address: "${REDIS_HOST}:6379" + use_tls: true + keys: + tia_livek_tia_api_key: "TiaLivekitSecret2026KeyBridgeSync01" \ No newline at end of file diff --git a/k8s/livekit/deployment.yaml b/k8s/livekit/deployment.yaml new file mode 100644 index 0000000..612c4f8 --- /dev/null +++ b/k8s/livekit/deployment.yaml @@ -0,0 +1,100 @@ +apiVersion: apps/v1 +kind: Deployment +metadata: + name: ${APP_NAME}-livekit + namespace: ${K8S_NAMESPACE} + labels: + app: ${APP_NAME}-livekit + app.kubernetes.io/name: tia-livekit + app.kubernetes.io/part-of: tia +spec: + replicas: ${LIVEKIT_REPLICAS} + strategy: + type: Recreate + selector: + matchLabels: + app: ${APP_NAME} + template: + metadata: + labels: + app: ${APP_NAME} + app.kubernetes.io/name: tia-livekit + app.kubernetes.io/part-of: tia + spec: + hostAliases: + - ip: ${REDIS_IP} + hostnames: + - ${REDIS_HOST} + securityContext: + runAsNonRoot: true + runAsUser: 1000 + runAsGroup: 1000 + fsGroup: 1000 + volumes: + - name: config-volume + configMap: + name: ${APP_NAME}-livekit-config + containers: + - name: ${APP_NAME}-livekit + image: ${IMAGE_REPOSITORY_LIVEKIT}:${IMAGE_TAG} + imagePullPolicy: IfNotPresent + args: + - "--config" + - "/etc/livekit/livekit.yaml" + envFrom: + - secretRef: + name: ${APP_NAME}-api-secrets + volumeMounts: + - name: config-volume + mountPath: /etc/livekit + readOnly: true + ports: + - name: signal + containerPort: 7880 + - name: rtc-tcp + containerPort: 7881 + - name: rtc-udp + containerPort: 7882 + protocol: UDP + readinessProbe: + tcpSocket: + port: 7880 + initialDelaySeconds: 20 + periodSeconds: 10 + timeoutSeconds: 5 + failureThreshold: 3 + livenessProbe: + tcpSocket: + port: 7880 + initialDelaySeconds: 30 + periodSeconds: 20 + timeoutSeconds: 5 + failureThreshold: 3 + resources: + requests: + cpu: "${CPU_LIVEKIT_REQ}" + memory: "${MEM_LIVEKIT_REQ}" + limits: + cpu: "${CPU_LIVEKIT_LIM}" + memory: "${MEM_LIVEKIT_LIM}" +--- +apiVersion: v1 +kind: Service +metadata: + name: ${APP_NAME}-livekit + namespace: ${K8S_NAMESPACE} +spec: + type: NodePort + selector: + app.kubernetes.io/name: tia-livekit + ports: + - name: signal + port: 7880 + targetPort: 7880 + - name: rtc-tcp + port: 7881 + targetPort: 7881 + - name: rtc-udp + port: 7882 + targetPort: 7882 + protocol: UDP diff --git a/k8s/regional/deployment-region.yaml b/k8s/regional/deployment-region.yaml new file mode 100644 index 0000000..15be7b7 --- /dev/null +++ b/k8s/regional/deployment-region.yaml @@ -0,0 +1,211 @@ +apiVersion: apps/v1 +kind: Deployment +metadata: + name: ${APP_NAME}-${REGION_ID} + namespace: ${K8S_NAMESPACE} + labels: + app: ${APP_NAME}-regional + tia-region: ${REGION_ID} +spec: + replicas: ${REGION_REPLICAS} + strategy: + type: RollingUpdate + rollingUpdate: + maxUnavailable: 0 + maxSurge: 1 + selector: + matchLabels: + app: ${APP_NAME}-regional + tia-region: ${REGION_ID} + template: + metadata: + labels: + app: ${APP_NAME}-regional + tia-region: ${REGION_ID} + spec: + terminationGracePeriodSeconds: ${TERMINATION_GRACE_SECONDS} + securityContext: + runAsNonRoot: true + runAsUser: 1000 + runAsGroup: 1000 + fsGroup: 1000 + containers: + - name: bridge + image: ${IMAGE_REPOSITORY}:${IMAGE_TAG} + imagePullPolicy: IfNotPresent + args: ["app.bridge_entry", "--host", "0.0.0.0", "--port", "8000", "--log-level", "info"] + ports: + - name: bridge-http + containerPort: 8000 + env: + - name: GOOGLE_APPLICATION_CREDENTIALS + value: /etc/google/credentials.json + - name: PYTHONPATH + value: /app/src + - name: REQUESTS_CA_BUNDLE + value: /etc/ssl/custom/tls.crt + - name: SSL_CERT_FILE + value: /etc/ssl/custom/tls.crt + - name: TIA_XAI_REGION + value: ${REGION_ID} + envFrom: + - configMapRef: + name: ${APP_NAME}-config + - secretRef: + name: ${APP_NAME}-api-secrets + readinessProbe: + httpGet: {path: /health, port: 8000} + initialDelaySeconds: 15 + periodSeconds: 5 + timeoutSeconds: 3 + failureThreshold: 3 + livenessProbe: + httpGet: {path: /health, port: 8000} + initialDelaySeconds: 30 + periodSeconds: 20 + timeoutSeconds: 5 + failureThreshold: 3 + resources: + requests: {cpu: "${CPU_TIA_BRIDGE_REQ}", memory: "${MEM_TIA_BRIDGE_REQ}"} + limits: {cpu: "${CPU_TIA_BRIDGE_LIM}", memory: "${MEM_TIA_BRIDGE_LIM}"} + volumeMounts: + - {name: google-sa-volume, mountPath: /etc/google, readOnly: true} + - {name: trusted-ca-volume, mountPath: /etc/ssl/custom, readOnly: true} + + - name: agent + image: ${IMAGE_REPOSITORY}:${IMAGE_TAG} + imagePullPolicy: IfNotPresent + args: ["app.agent_entry", "start", "--log-level", "info"] + ports: + - name: agent-http + containerPort: 18081 + envFrom: + - configMapRef: + name: ${APP_NAME}-config + - secretRef: + name: ${APP_NAME}-api-secrets + env: + - name: GOOGLE_APPLICATION_CREDENTIALS + value: /etc/google/credentials.json + - name: PYTHONPATH + value: /app/src + - name: AGENT_SERVER_PORT + value: "18081" + - name: NUM_IDLE_PROCESSES + value: "1" + - name: REQUESTS_CA_BUNDLE + value: /etc/ssl/custom/tls.crt + - name: SSL_CERT_FILE + value: /etc/ssl/custom/tls.crt + - name: TIA_XAI_REGION + value: ${REGION_ID} + # Agent sees a local xAI-compatible endpoint. Real OCI credentials stay in xai-pool. + - name: XAI_WEBSOCKET_URL + value: ws://127.0.0.1:18100/xai/v1/tts + - name: XAI_TTS_AUTH_METHOD + value: API_KEY + - name: XAI_API_KEY + value: local-pool-proxy + startupProbe: + httpGet: {path: /, port: 18081} + initialDelaySeconds: 10 + periodSeconds: 5 + timeoutSeconds: 5 + failureThreshold: 24 + readinessProbe: + httpGet: {path: /, port: 18081} + initialDelaySeconds: 20 + periodSeconds: 10 + timeoutSeconds: 5 + failureThreshold: 3 + livenessProbe: + httpGet: {path: /, port: 18081} + initialDelaySeconds: 30 + periodSeconds: 20 + timeoutSeconds: 5 + failureThreshold: 3 + resources: + requests: {cpu: "${CPU_TIA_REQ}", memory: "${MEM_TIA_REQ}"} + limits: {cpu: "${CPU_TIA_LIM}", memory: "${MEM_TIA_LIM}"} + volumeMounts: + - {name: google-sa-volume, mountPath: /etc/google, readOnly: true} + - {name: trusted-ca-volume, mountPath: /etc/ssl/custom, readOnly: true} + + - name: xai-pool + image: ${IMAGE_REPOSITORY}:${IMAGE_TAG} + imagePullPolicy: IfNotPresent + args: ["app.livekit.adapters.xai_pool_proxy"] + ports: + - name: xai-pool + containerPort: 18100 + env: + - name: PYTHONPATH + value: /app/src + - name: TIA_XAI_REGION + value: ${REGION_ID} + - name: XAI_POOL_UPSTREAM_URL + value: ${XAI_UPSTREAM_URL} + - name: XAI_POOL_SIZE + value: "${XAI_POOL_SIZE}" + - name: XAI_POOL_UNAVAILABLE_FREE + value: "${XAI_POOL_UNAVAILABLE_FREE}" + - name: XAI_POOL_RECOVER_FREE + value: "${XAI_POOL_RECOVER_FREE}" + - name: XAI_POOL_CONNECTION_TTL_S + value: "${XAI_POOL_CONNECTION_TTL_S}" + - name: XAI_POOL_REFRESH_JITTER_S + value: "${XAI_POOL_REFRESH_JITTER_S}" + - name: XAI_POOL_PREWARM_CONCURRENCY + value: "${XAI_POOL_PREWARM_CONCURRENCY}" + - name: XAI_TTS_VOICE + value: ${XAI_TTS_VOICE} + - name: XAI_TTS_LANGUAGE + value: ${XAI_TTS_LANGUAGE} + - name: XAI_TTS_AUTH_METHOD + value: ${XAI_UPSTREAM_AUTH_METHOD} + - name: OCI_COMPARTMENT_ID + valueFrom: + secretKeyRef: + name: ${XAI_SECRET_NAME} + key: OCI_COMPARTMENT_ID + optional: true + - name: XAI_API_KEY + valueFrom: + secretKeyRef: + name: ${XAI_SECRET_NAME} + key: XAI_API_KEY + optional: true + - name: REQUESTS_CA_BUNDLE + value: /etc/ssl/custom/tls.crt + - name: SSL_CERT_FILE + value: /etc/ssl/custom/tls.crt + readinessProbe: + httpGet: {path: /readyz, port: 18100} + initialDelaySeconds: 5 + periodSeconds: 2 + timeoutSeconds: 1 + failureThreshold: 2 + successThreshold: 1 + livenessProbe: + httpGet: {path: /healthz, port: 18100} + initialDelaySeconds: 10 + periodSeconds: 10 + timeoutSeconds: 2 + failureThreshold: 3 + lifecycle: + preStop: + exec: + command: ["/bin/sh", "-c", "curl -sf -X POST http://127.0.0.1:18100/drain || true; sleep ${DRAIN_SECONDS}"] + resources: + requests: {cpu: "${CPU_XAI_POOL_REQ}", memory: "${MEM_XAI_POOL_REQ}"} + limits: {cpu: "${CPU_XAI_POOL_LIM}", memory: "${MEM_XAI_POOL_LIM}"} + volumeMounts: + - {name: trusted-ca-volume, mountPath: /etc/ssl/custom, readOnly: true} + + volumes: + - name: google-sa-volume + secret: + secretName: ${APP_NAME}-google-sa-secret + - name: trusted-ca-volume + secret: + secretName: shared-tls-secret diff --git a/k8s/regional/hpa-region.yaml b/k8s/regional/hpa-region.yaml new file mode 100644 index 0000000..2215cfe --- /dev/null +++ b/k8s/regional/hpa-region.yaml @@ -0,0 +1,32 @@ +apiVersion: autoscaling/v2 +kind: HorizontalPodAutoscaler +metadata: + name: ${APP_NAME}-${REGION_ID} + namespace: ${K8S_NAMESPACE} +spec: + scaleTargetRef: + apiVersion: apps/v1 + kind: Deployment + name: ${APP_NAME}-${REGION_ID} + minReplicas: ${HPA_MIN_REPLICAS} + maxReplicas: ${HPA_MAX_REPLICAS} + behavior: + scaleUp: + stabilizationWindowSeconds: 0 + policies: + - type: Percent + value: 100 + periodSeconds: 60 + scaleDown: + stabilizationWindowSeconds: 300 + policies: + - type: Percent + value: 25 + periodSeconds: 60 + metrics: + - type: Resource + resource: + name: cpu + target: + type: Utilization + averageUtilization: ${HPA_CPU_TARGET} diff --git a/k8s/regional/pdb.yaml b/k8s/regional/pdb.yaml new file mode 100644 index 0000000..eac869e --- /dev/null +++ b/k8s/regional/pdb.yaml @@ -0,0 +1,10 @@ +apiVersion: policy/v1 +kind: PodDisruptionBudget +metadata: + name: ${APP_NAME}-regional + namespace: ${K8S_NAMESPACE} +spec: + minAvailable: ${PDB_MIN_AVAILABLE} + selector: + matchLabels: + app: ${APP_NAME}-regional diff --git a/k8s/regional/regions.env.example b/k8s/regional/regions.env.example new file mode 100644 index 0000000..a471928 --- /dev/null +++ b/k8s/regional/regions.env.example @@ -0,0 +1,51 @@ +# Common +APP_NAME=tim-ai-atend-agnt-integ-tia +K8S_NAMESPACE=agnt-ai-atendimento +IMAGE_REPOSITORY=iad.ocir.io/SEU_NAMESPACE/tia +IMAGE_TAG=regional-xai-pool-v1 +TIA_SERVICE_TYPE=LoadBalancer +TERMINATION_GRACE_SECONDS=600 +DRAIN_SECONDS=30 +PDB_MIN_AVAILABLE=2 + +# Existing TIA resources +CPU_TIA_BRIDGE_REQ=250m +MEM_TIA_BRIDGE_REQ=512Mi +CPU_TIA_BRIDGE_LIM=1000m +MEM_TIA_BRIDGE_LIM=1Gi +CPU_TIA_REQ=500m +MEM_TIA_REQ=1Gi +CPU_TIA_LIM=2000m +MEM_TIA_LIM=2Gi +CPU_XAI_POOL_REQ=200m +MEM_XAI_POOL_REQ=256Mi +CPU_XAI_POOL_LIM=1000m +MEM_XAI_POOL_LIM=768Mi + +# Pool profile - 50 means 50 prewarmed upstream WebSockets PER POD. +XAI_POOL_SIZE=50 +XAI_POOL_UNAVAILABLE_FREE=2 +XAI_POOL_RECOVER_FREE=5 +XAI_POOL_CONNECTION_TTL_S=540 +XAI_POOL_REFRESH_JITTER_S=45 +XAI_POOL_PREWARM_CONCURRENCY=5 +XAI_TTS_VOICE=c8x2ieiocufs +XAI_TTS_LANGUAGE=pt-BR +XAI_UPSTREAM_AUTH_METHOD=API_KEY + +# HPA. WARNING: replicas * XAI_POOL_SIZE must respect OCI/xAI quota. +HPA_MIN_REPLICAS=1 +HPA_MAX_REPLICAS=3 +HPA_CPU_TARGET=65 + +# ORD +ORD_REGION_ID=ord +ORD_REGION_REPLICAS=1 +ORD_XAI_UPSTREAM_URL=wss://peordagnt002prd.pe.inference.generativeai.us-chicago-1.oci.oraclecloud.com/xai/v1/tts +ORD_XAI_SECRET_NAME=xai-ord-credentials + +# IAD +IAD_REGION_ID=iad +IAD_REGION_REPLICAS=1 +IAD_XAI_UPSTREAM_URL=wss://peiadagnt003prd.pe.inference.generativeai.us-ashburn-1.oci.oraclecloud.com/xai/v1/tts +IAD_XAI_SECRET_NAME=xai-iad-credentials diff --git a/k8s/regional/service.yaml b/k8s/regional/service.yaml new file mode 100644 index 0000000..54f2891 --- /dev/null +++ b/k8s/regional/service.yaml @@ -0,0 +1,17 @@ +apiVersion: v1 +kind: Service +metadata: + name: ${APP_NAME}-regional + namespace: ${K8S_NAMESPACE} + labels: + app: ${APP_NAME}-regional +spec: + type: ${TIA_SERVICE_TYPE} + sessionAffinity: None + selector: + app: ${APP_NAME}-regional + ports: + - name: ws-http + protocol: TCP + port: 80 + targetPort: 8000 diff --git a/k8s/regional/xai-secret.example.yaml b/k8s/regional/xai-secret.example.yaml new file mode 100644 index 0000000..8102ff5 --- /dev/null +++ b/k8s/regional/xai-secret.example.yaml @@ -0,0 +1,19 @@ +apiVersion: v1 +kind: Secret +metadata: + name: xai-ord-credentials + namespace: ${K8S_NAMESPACE} +type: Opaque +stringData: + XAI_API_KEY: "REPLACE_ME" + OCI_COMPARTMENT_ID: "" +--- +apiVersion: v1 +kind: Secret +metadata: + name: xai-iad-credentials + namespace: ${K8S_NAMESPACE} +type: Opaque +stringData: + XAI_API_KEY: "REPLACE_ME" + OCI_COMPARTMENT_ID: "" diff --git a/k8s/secrets.yaml b/k8s/secrets.yaml new file mode 100644 index 0000000..8c4b3d5 --- /dev/null +++ b/k8s/secrets.yaml @@ -0,0 +1,11 @@ +apiVersion: v1 +kind: Secret +metadata: + name: ${APP_NAME}-api-secrets + namespace: ${K8S_NAMESPACE} +type: Opaque +stringData: + AZURE_SPEECH_KEY: "${AZURE_SPEECH_KEY}" + XAI_API_KEY: "${XAI_API_KEY}" + LIVEKIT_REDIS_USERNAME: "${LIVEKIT_REDIS_USERNAME}" + LIVEKIT_REDIS_PASSWORD: "${LIVEKIT_REDIS_PASSWORD}" \ No newline at end of file diff --git a/k8s/tia/Dockerfile b/k8s/tia/Dockerfile new file mode 100644 index 0000000..73de5ce --- /dev/null +++ b/k8s/tia/Dockerfile @@ -0,0 +1,31 @@ +FROM python:3.12-slim + +WORKDIR /app + +ENV PYTHONDONTWRITEBYTECODE=1 +ENV PYTHONUNBUFFERED=1 +ENV PYTHONPATH=/app/src +ENV HF_HOME=/hf +ENV HF_HUB_CACHE=/hf/hub +ENV HUGGINGFACE_HUB_CACHE=/hf/hub +ENV TRANSFORMERS_CACHE=/hf/hub + +RUN apt-get update && apt-get install -y --no-install-recommends \ + build-essential \ + curl \ + libsndfile1 \ +&& rm -rf /var/lib/apt/lists/* + +COPY requirements.txt ./ +RUN pip install --no-cache-dir -r requirements.txt + +COPY src/ ./src/ +RUN mkdir -p "$HF_HUB_CACHE" && \ + python -m app.livekit.main download-files + +RUN useradd -m -u 1000 agent && chown -R agent:agent /app /hf +USER agent + +EXPOSE 8000 18081 + +ENTRYPOINT ["python", "-m"] diff --git a/k8s/tia/deployment.yaml b/k8s/tia/deployment.yaml new file mode 100644 index 0000000..bd69658 --- /dev/null +++ b/k8s/tia/deployment.yaml @@ -0,0 +1,264 @@ +apiVersion: apps/v1 +kind: Deployment +metadata: + name: ${APP_NAME}-app + namespace: ${K8S_NAMESPACE} +spec: + replicas: ${TIA_REPLICAS} + strategy: + type: Recreate + selector: + matchLabels: + app: ${APP_NAME}-app + template: + metadata: + labels: + app: ${APP_NAME}-app + spec: + hostAliases: + - ip: "10.151.225.135" + hostnames: + - "speech-agent-ai-atendi-fqa-01.cognitiveservices.azure.com" + - ip: 10.153.35.23 + hostnames: + - tim-ai-atend-agnt-opentelemetry + - ip: 10.154.0.154 + hostnames: + - peordagnt002prd.pe.inference.generativeai.us-chicago-1.oci.oraclecloud.com + - ip: 10.154.16.244 + hostnames: + - peiadagnt003prd.pe.inference.generativeai.us-ashburn-1.oci.oraclecloud.com + - ip: ${REDIS_IP} + hostnames: + - ${REDIS_HOST} + - ip: 10.154.16.244 + hostnames: + - peiadagnt003prd.pe.inference.generativeai.us-ashburn-1.oci.oraclecloud.com + securityContext: + runAsNonRoot: true + runAsUser: 1000 + runAsGroup: 1000 + fsGroup: 1000 + containers: + - name: ${APP_NAME}-app-bridge + image: ${IMAGE_REPOSITORY}:${IMAGE_TAG} + imagePullPolicy: IfNotPresent + args: + - app.bridge_entry + - --host + - 0.0.0.0 + - --port + - "8000" + - --log-level + - info + ports: + - name: http + containerPort: 8000 + env: + - name: GOOGLE_APPLICATION_CREDENTIALS + value: "/etc/google/credentials.json" + - name: PYTHONPATH + value: /app/src + - name: REQUESTS_CA_BUNDLE + value: "/etc/ssl/custom/tls.crt" + - name: SSL_CERT_FILE + value: "/etc/ssl/custom/tls.crt" + envFrom: + - configMapRef: + name: ${APP_NAME}-config + - secretRef: + name: ${APP_NAME}-api-secrets + volumeMounts: + - name: google-sa-volume + mountPath: /etc/google + readOnly: true + - name: trusted-ca-volume + mountPath: "/etc/ssl/custom" + readOnly: true + readinessProbe: + httpGet: + path: /health + port: 8000 + initialDelaySeconds: 20 + periodSeconds: 10 + timeoutSeconds: 5 + successThreshold: 1 + failureThreshold: 3 + livenessProbe: + httpGet: + path: /health + port: 8000 + initialDelaySeconds: 30 + periodSeconds: 20 + timeoutSeconds: 5 + failureThreshold: 3 + resources: + requests: + cpu: "${CPU_TIA_BRIDGE_REQ}" + memory: "${MEM_TIA_BRIDGE_REQ}" + limits: + cpu: "${CPU_TIA_BRIDGE_LIM}" + memory: "${MEM_TIA_BRIDGE_LIM}" + + - name: ${APP_NAME}-app-agent + image: ${IMAGE_REPOSITORY}:${IMAGE_TAG} + imagePullPolicy: IfNotPresent + args: + - app.agent_entry + - start + - --log-level + - info + ports: + - name: agent-http + containerPort: 18081 + env: + - name: GOOGLE_APPLICATION_CREDENTIALS + value: "/etc/google/credentials.json" + - name: PYTHONPATH + value: /app/src + - name: AGENT_SERVER_PORT + value: "18081" + - name: NUM_IDLE_PROCESSES + value: "1" + - name: REQUESTS_CA_BUNDLE + value: "/etc/ssl/custom/tls.crt" + - name: SSL_CERT_FILE + value: "/etc/ssl/custom/tls.crt" + envFrom: + - configMapRef: + name: ${APP_NAME}-config + - secretRef: + name: ${APP_NAME}-api-secrets + volumeMounts: + - name: google-sa-volume + mountPath: /etc/google + readOnly: true + - name: trusted-ca-volume + mountPath: "/etc/ssl/custom" + readOnly: true + startupProbe: + httpGet: + path: / + port: 18081 + initialDelaySeconds: 10 + periodSeconds: 5 + timeoutSeconds: 5 + failureThreshold: 24 + readinessProbe: + httpGet: + path: / + port: 18081 + initialDelaySeconds: 20 + periodSeconds: 10 + timeoutSeconds: 5 + successThreshold: 1 + failureThreshold: 3 + livenessProbe: + httpGet: + path: / + port: 18081 + initialDelaySeconds: 30 + periodSeconds: 20 + timeoutSeconds: 5 + failureThreshold: 3 + resources: + requests: + cpu: "${CPU_TIA_REQ}" + memory: "${MEM_TIA_REQ}" + limits: + cpu: "${CPU_TIA_LIM}" + memory: "${MEM_TIA_LIM}" + volumes: + - name: google-sa-volume + secret: + secretName: ${APP_NAME}-google-sa-secret + - name: trusted-ca-volume + secret: + secretName: shared-tls-secret +--- +apiVersion: v1 +kind: Service +metadata: + name: ${APP_NAME}-app + namespace: ${K8S_NAMESPACE} + labels: + app: ${APP_NAME}-app +spec: + type: NodePort + selector: + app: ${APP_NAME}-app + ports: + - name: http + protocol: TCP + port: 80 + targetPort: 8000 + - name: https + protocol: TCP + port: 443 + targetPort: 8000 +--- +apiVersion: gateway.networking.k8s.io/v1 +kind: HTTPRoute +metadata: + name: ${APP_NAME}-route + namespace: ${K8S_NAMESPACE} +spec: + parentRefs: + - name: istio-gateway + namespace: istio-gateway + hostnames: + - ${APP_NAME} + rules: + - matches: + - path: + type: PathPrefix + value: / + timeouts: + request: 1600s + backendRefs: + - name: ${APP_NAME}-app + port: 80 +--- +apiVersion: gateway.networking.k8s.io/v1 +kind: HTTPRoute +metadata: + name: ${APP_NAME}-route-http + namespace: ${K8S_NAMESPACE} +spec: + parentRefs: + - name: istio-gateway + namespace: istio-gateway + hostnames: + - ${DNS} + rules: + - matches: + - path: + type: PathPrefix + value: / + timeouts: + request: 1600s + backendRefs: + - name: ${APP_NAME}-app + port: 80 +--- +apiVersion: gateway.networking.k8s.io/v1 +kind: HTTPRoute +metadata: + name: ${APP_NAME}-route-https + namespace: ${K8S_NAMESPACE} +spec: + parentRefs: + - name: istio-gateway + namespace: istio-gateway + hostnames: + - ${DNS} + rules: + - matches: + - path: + type: PathPrefix + value: / + timeouts: + request: 1600s + backendRefs: + - name: ${APP_NAME}-app + port: 443 diff --git a/livekit.yaml b/livekit.yaml new file mode 100644 index 0000000..7501d0d --- /dev/null +++ b/livekit.yaml @@ -0,0 +1,8 @@ +port: 7880 +log_level: warn +rtc: + tcp_port: 7881 + port_range_start: 50000 + port_range_end: 60000 +keys: + tia_livek_tia_api_key: "TiaLivekitSecret2026KeyBridgeSync01" diff --git a/makefile b/makefile new file mode 100644 index 0000000..7e3119b --- /dev/null +++ b/makefile @@ -0,0 +1,495 @@ +SHELL := /bin/bash +OS := $(shell uname -s) + +# ====================== +# Local (DEV) - python +# ====================== +VENV ?= .venv +PY ?= $(VENV)/bin/python +BOOTSTRAP_PY ?= +SRC ?= app +PYTHONPATH_HOST ?= $(CURDIR)/src +PYTHONPATH_CONT ?= /app/src + +AGENT_FILE ?= $(SRC).agent_entry +BRIDGE_MODULE ?= $(SRC).bridge_entry +RUN_DIR ?= .run +AGENT_PID_FILE ?= $(RUN_DIR)/agent.pid +BRIDGE_PID_FILE ?= $(RUN_DIR)/bridge.pid +AGENT_LOG_FILE ?= $(RUN_DIR)/agent.log +BRIDGE_LOG_FILE ?= $(RUN_DIR)/bridge.log + +WS_HOST ?= 0.0.0.0 +DEV_WS_PORT ?= 8000 +AGENT_SERVER_PORT ?= 18081 + +# envs +ENV_DEV ?= .env.dev +ENV_PROD ?= .env.prod + +# ====================== +# Podman (PROD) - host network +# ====================== +IMAGE ?= api-tia-lk:latest +POD_NAME ?= api-tia-pod + +PROD_WS_PORT ?= 8000 +EXPORT_DIR_HOST ?= /RemoteChannelsOutBound +EXPORT_DIR_CONT ?= /RemoteChannelsOutBound + +# HuggingFace cache (turn-detector models) +HF_CACHE_VOL ?= lk-hf-cache +HF_HOME_CONT ?= /hf +HF_HUB_CACHE_CONT ?= /hf/hub + +# ====================== +# LiveKit (build/run) - host network +# ====================== +LIVEKIT_IMAGE ?= livekit/livekit-server:latest +LIVEKIT_NAME ?= livekit-prod +LIVEKIT_CFG ?= $(CURDIR)/livekit.yaml +LIVEKIT_PORT_HTTP ?= 7880 +LIVEKIT_PORT_TCP ?= 7881 + +.PHONY: help \ + setup check-venv \ + agent agent-models bridge test \ + agent-up agent-down agent-status agent-logs \ + bridge-up bridge-down bridge-status bridge-logs \ + local-up local-down local-status local-logs local-stresstest \ + docker-build docker-models docker-up docker-down docker-status docker-logs docker-logs-bridge \ + livekit livekit-stop livekit-logs livekit-status livekit-config-check \ + copy-logs download-call-segments download-entire-calls + +help: + @echo "Targets (DEV local):" + @echo " make setup - cria $(VENV), instala dependencias e prepara $(ENV_DEV)" + @echo " make local-up - sobe livekit + bridge + agent em background" + @echo " make local-down - derruba bridge + agent locais e para o livekit" + @echo " make local-status - mostra status do livekit + pids locais" + @echo " make local-logs - informa onde estao os logs locais" + @echo " make local-stresstest - valida STT/TTS/Bridge/LiveKit reais e gera report.md" + @echo " make agent - roda agent (usa $(ENV_DEV))" + @echo " make agent-up - sobe agent em background" + @echo " make agent-down - derruba agent local em background" + @echo " make agent-logs - tail do log local do agent" + @echo " make agent-models - baixa modelos locais do turn-detector (usa $(ENV_DEV))" + @echo " make bridge - roda bridge (usa $(ENV_DEV)) na porta $(DEV_WS_PORT)" + @echo " make bridge-up - sobe bridge em background" + @echo " make bridge-down - derruba bridge local em background" + @echo " make bridge-logs - tail do log local do bridge" + @echo " make download-call-segments DATE=AAAA-MM-DD SESSION_ID=id [BUCKET=nome]" + @echo " make download-entire-calls DATE=AAAA-MM-DD [BUCKET=nome]" + @echo " make test - roda a suite pytest com PYTHONPATH=src" + @echo + @echo "Targets (PROD podman / host network):" + @echo " make docker-build - build imagem (usa PROD por padrão)" + @echo " make docker-models - baixa modelos do turn-detector no volume $(HF_CACHE_VOL)" + @echo " make docker-up - sobe agent+bridge no pod (host network), bridge na porta $(PROD_WS_PORT)" + @echo " make docker-logs - logs do agent" + @echo " make docker-logs-bridge - logs do bridge" + @echo " make docker-down - derruba pod" + @echo + @echo "Targets (LiveKit host network):" + @echo " make livekit - sobe LiveKit (host network)" + @echo " make livekit-logs - logs LiveKit" + @echo " make livekit-stop - remove container LiveKit" + +# ====================== +# DEV (local) +# ====================== +setup: + @if [ -x "$(PY)" ]; then \ + if $(PY) -c 'import sys; raise SystemExit(0 if (sys.version_info.major == 3 and 9 <= sys.version_info.minor <= 13) else 1)' >/dev/null 2>&1; then \ + echo "[make] usando virtualenv existente em $(VENV)"; \ + else \ + echo "ERRO: $(VENV) usa uma versao de Python nao suportada por este projeto."; \ + echo "Remova $(VENV) e rode 'make BOOTSTRAP_PY=python3.12 setup'."; \ + exit 1; \ + fi; \ + else \ + bootstrap_py="$(BOOTSTRAP_PY)"; \ + if [ -n "$$bootstrap_py" ]; then \ + if ! command -v "$$bootstrap_py" >/dev/null 2>&1; then \ + echo "ERRO: interpretador '$$bootstrap_py' nao encontrado no PATH."; \ + exit 1; \ + fi; \ + if ! "$$bootstrap_py" -c 'import sys; raise SystemExit(0 if (sys.version_info.major == 3 and 9 <= sys.version_info.minor <= 13) else 1)' >/dev/null 2>&1; then \ + echo "ERRO: '$$bootstrap_py' nao eh compativel. Use Python 3.9-3.13."; \ + exit 1; \ + fi; \ + else \ + for candidate in python3.13 python3.12 python3.11 python3.10 python3.9 python3 python; do \ + if command -v "$$candidate" >/dev/null 2>&1 && "$$candidate" -c 'import sys; raise SystemExit(0 if (sys.version_info.major == 3 and 9 <= sys.version_info.minor <= 13) else 1)' >/dev/null 2>&1; then \ + bootstrap_py="$$candidate"; \ + break; \ + fi; \ + done; \ + fi; \ + if [ -n "$$bootstrap_py" ]; then \ + echo "[make] criando virtualenv com $$bootstrap_py"; \ + "$$bootstrap_py" -m venv "$(VENV)"; \ + elif command -v uv >/dev/null 2>&1; then \ + echo "ERRO: nenhum Python compativel encontrado no PATH."; \ + echo "Instale Python 3.9-3.13 ou rode 'make BOOTSTRAP_PY=python3.12 setup'."; \ + exit 1; \ + else \ + echo "ERRO: nenhum interpretador encontrado. Instale Python 3.9-3.13 e rode 'make setup' novamente."; \ + exit 1; \ + fi; \ + fi + @$(PY) -m pip install --upgrade pip + @$(PY) -m pip install -r requirements.txt + @if [ ! -f "$(ENV_DEV)" ]; then \ + cp .env.example "$(ENV_DEV)"; \ + echo "[make] criado $(ENV_DEV) a partir de .env.example"; \ + fi + @echo "[make] ambiente pronto. Para rodar os testes: make test" + +check-venv: + @if [ ! -x "$(PY)" ]; then \ + echo "ERRO: virtualenv nao encontrado em $(PY). Rode 'make setup' antes."; \ + exit 1; \ + fi + @if ! $(PY) -c 'import sys; raise SystemExit(0 if (sys.version_info.major == 3 and 9 <= sys.version_info.minor <= 13) else 1)' >/dev/null 2>&1; then \ + echo "ERRO: o virtualenv em $(VENV) nao usa um Python suportado."; \ + echo "Remova $(VENV) e rode 'make BOOTSTRAP_PY=python3.12 setup'."; \ + exit 1; \ + fi + +agent: check-venv + @set -a; [ -f "$(ENV_DEV)" ] && source "$(ENV_DEV)"; set +a; \ + AGENT_SERVER_PORT="$(AGENT_SERVER_PORT)" \ + STT_DUMP_DIR="$${AGENT_DIAG_STT_DUMP_DIR:-$${STT_DUMP_DIR:-}}" \ + FLOW_LOG_VAD_DECISIONS="$${AGENT_DIAG_FLOW_LOG_VAD_DECISIONS:-$${FLOW_LOG_VAD_DECISIONS:-0}}" \ + FLOW_LOG_VAD_ACTIVITY="$${AGENT_DIAG_FLOW_LOG_VAD_ACTIVITY:-$${FLOW_LOG_VAD_ACTIVITY:-0}}" \ + FLOW_LOG_VAD_ACTIVITY_MIN_PROB="$${AGENT_DIAG_FLOW_LOG_VAD_ACTIVITY_MIN_PROB:-$${FLOW_LOG_VAD_ACTIVITY_MIN_PROB:-0.03}}" \ + PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" \ + $(PY) -m $(AGENT_FILE) start --log-level info + +agent-models: check-venv + @set -a; [ -f "$(ENV_DEV)" ] && source "$(ENV_DEV)"; set +a; \ + AGENT_SERVER_PORT="$(AGENT_SERVER_PORT)" PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" $(PY) -m $(AGENT_FILE) download-files + +bridge: check-venv + @set -a; [ -f "$(ENV_DEV)" ] && source "$(ENV_DEV)"; set +a; \ + PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" $(PY) -m $(BRIDGE_MODULE) --host $(WS_HOST) --port $(DEV_WS_PORT) --log-level info + +download-call-segments: check-venv + @test -n "$(DATE)" || (echo "ERRO: informe DATE=AAAA-MM-DD"; exit 2) + @test -n "$(SESSION_ID)" || (echo "ERRO: informe SESSION_ID="; exit 2) + @set -a; [ -f "$(ENV_DEV)" ] && source "$(ENV_DEV)"; set +a; \ + PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" \ + $(PY) -m app.tools.oci_audio_download \ + $(if $(BUCKET),--bucket "$(BUCKET)",) \ + segments --date "$(DATE)" --session-id "$(SESSION_ID)" + +download-entire-calls: check-venv + @test -n "$(DATE)" || (echo "ERRO: informe DATE=AAAA-MM-DD"; exit 2) + @set -a; [ -f "$(ENV_DEV)" ] && source "$(ENV_DEV)"; set +a; \ + PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" \ + $(PY) -m app.tools.oci_audio_download \ + $(if $(BUCKET),--bucket "$(BUCKET)",) \ + entire-calls --date "$(DATE)" + +test: check-venv + PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" $(PY) -m pytest + +agent-up: check-venv + @mkdir -p "$(RUN_DIR)" + @if [ -f "$(AGENT_PID_FILE)" ] && kill -0 "$$(cat "$(AGENT_PID_FILE)")" >/dev/null 2>&1; then \ + echo "[make] agent ja esta rodando com pid $$(cat "$(AGENT_PID_FILE)")"; \ + exit 0; \ + fi + @port_pids="$$(lsof -t -iTCP:$(AGENT_SERVER_PORT) -sTCP:LISTEN 2>/dev/null || true)"; \ + if [ -n "$$port_pids" ]; then \ + echo "ERRO: a porta $(AGENT_SERVER_PORT) ja esta em uso pelos pid(s): $$port_pids"; \ + echo "Rode 'make agent-down' para limpar o worker local antigo ou finalize o processo manualmente."; \ + exit 1; \ + fi + @nohup /bin/bash -lc 'set -a; [ -f "$(ENV_DEV)" ] && source "$(ENV_DEV)"; set +a; AGENT_SERVER_PORT="$(AGENT_SERVER_PORT)" STT_DUMP_DIR="$${AGENT_DIAG_STT_DUMP_DIR:-$${STT_DUMP_DIR:-}}" FLOW_LOG_VAD_DECISIONS="$${AGENT_DIAG_FLOW_LOG_VAD_DECISIONS:-$${FLOW_LOG_VAD_DECISIONS:-0}}" FLOW_LOG_VAD_ACTIVITY="$${AGENT_DIAG_FLOW_LOG_VAD_ACTIVITY:-$${FLOW_LOG_VAD_ACTIVITY:-0}}" FLOW_LOG_VAD_ACTIVITY_MIN_PROB="$${AGENT_DIAG_FLOW_LOG_VAD_ACTIVITY_MIN_PROB:-$${FLOW_LOG_VAD_ACTIVITY_MIN_PROB:-0.03}}" PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" exec "$(PY)" -m $(AGENT_FILE) start --log-level info' >"$(AGENT_LOG_FILE)" 2>&1 & echo $$! >"$(AGENT_PID_FILE)" + @sleep 1 + @if kill -0 "$$(cat "$(AGENT_PID_FILE)")" >/dev/null 2>&1; then \ + echo "[make] agent iniciado em background (pid $$(cat "$(AGENT_PID_FILE)"))"; \ + echo "[make] log: $(AGENT_LOG_FILE)"; \ + else \ + echo "ERRO: agent falhou ao iniciar. Veja $(AGENT_LOG_FILE)"; \ + rm -f "$(AGENT_PID_FILE)"; \ + exit 1; \ + fi + +agent-down: + @if [ -f "$(AGENT_PID_FILE)" ]; then \ + pid="$$(cat "$(AGENT_PID_FILE)")"; \ + if kill -0 "$$pid" >/dev/null 2>&1; then \ + kill "$$pid" >/dev/null 2>&1 || true; \ + sleep 1; \ + if kill -0 "$$pid" >/dev/null 2>&1; then \ + kill -9 "$$pid" >/dev/null 2>&1 || true; \ + echo "[make] agent encerrado com SIGKILL (pid $$pid)"; \ + else \ + echo "[make] agent encerrado (pid $$pid)"; \ + fi; \ + else \ + echo "[make] agent nao estava rodando, removendo pid stale"; \ + fi; \ + rm -f "$(AGENT_PID_FILE)"; \ + else \ + echo "[make] agent nao esta rodando"; \ + fi + @port_pids="$$(lsof -t -iTCP:$(AGENT_SERVER_PORT) -sTCP:LISTEN 2>/dev/null || true)"; \ + if [ -n "$$port_pids" ]; then \ + for pid in $$port_pids; do \ + kill "$$pid" >/dev/null 2>&1 || true; \ + sleep 1; \ + if kill -0 "$$pid" >/dev/null 2>&1; then \ + kill -9 "$$pid" >/dev/null 2>&1 || true; \ + echo "[make] worker do agent encerrado com SIGKILL na porta $(AGENT_SERVER_PORT) (pid $$pid)"; \ + else \ + echo "[make] worker do agent encerrado na porta $(AGENT_SERVER_PORT) (pid $$pid)"; \ + fi; \ + done; \ + fi + +agent-status: + @if [ -f "$(AGENT_PID_FILE)" ] && kill -0 "$$(cat "$(AGENT_PID_FILE)")" >/dev/null 2>&1; then \ + echo "agent: running (pid $$(cat "$(AGENT_PID_FILE)"))"; \ + elif [ -n "$$(lsof -t -iTCP:$(AGENT_SERVER_PORT) -sTCP:LISTEN 2>/dev/null || true)" ]; then \ + echo "agent: running (listener na porta $(AGENT_SERVER_PORT) sem pid file)"; \ + else \ + echo "agent: stopped"; \ + fi + +agent-logs: + @if [ ! -f "$(AGENT_LOG_FILE)" ]; then \ + echo "ERRO: log do agent nao encontrado em $(AGENT_LOG_FILE)"; \ + exit 1; \ + fi + @tail -f "$(AGENT_LOG_FILE)" + +bridge-up: check-venv + @mkdir -p "$(RUN_DIR)" + @if [ -f "$(BRIDGE_PID_FILE)" ] && kill -0 "$$(cat "$(BRIDGE_PID_FILE)")" >/dev/null 2>&1; then \ + echo "[make] bridge ja esta rodando com pid $$(cat "$(BRIDGE_PID_FILE)")"; \ + exit 0; \ + fi + @port_pids="$$(lsof -t -iTCP:$(DEV_WS_PORT) -sTCP:LISTEN 2>/dev/null || true)"; \ + if [ -n "$$port_pids" ]; then \ + echo "ERRO: a porta $(DEV_WS_PORT) ja esta em uso pelos pid(s): $$port_pids"; \ + echo "Rode 'make bridge-down' para limpar o bridge local antigo ou finalize o processo manualmente."; \ + exit 1; \ + fi + @nohup /bin/bash -lc 'set -a; [ -f "$(ENV_DEV)" ] && source "$(ENV_DEV)"; set +a; PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" exec "$(PY)" -m $(BRIDGE_MODULE) --host "$(WS_HOST)" --port "$(DEV_WS_PORT)" --log-level info' >"$(BRIDGE_LOG_FILE)" 2>&1 & echo $$! >"$(BRIDGE_PID_FILE)" + @sleep 1 + @if kill -0 "$$(cat "$(BRIDGE_PID_FILE)")" >/dev/null 2>&1; then \ + echo "[make] bridge iniciado em background (pid $$(cat "$(BRIDGE_PID_FILE)"))"; \ + echo "[make] log: $(BRIDGE_LOG_FILE)"; \ + else \ + echo "ERRO: bridge falhou ao iniciar. Veja $(BRIDGE_LOG_FILE)"; \ + rm -f "$(BRIDGE_PID_FILE)"; \ + exit 1; \ + fi + +bridge-down: + @if [ -f "$(BRIDGE_PID_FILE)" ]; then \ + pid="$$(cat "$(BRIDGE_PID_FILE)")"; \ + if kill -0 "$$pid" >/dev/null 2>&1; then \ + kill "$$pid" >/dev/null 2>&1 || true; \ + sleep 1; \ + if kill -0 "$$pid" >/dev/null 2>&1; then \ + kill -9 "$$pid" >/dev/null 2>&1 || true; \ + echo "[make] bridge encerrado com SIGKILL (pid $$pid)"; \ + else \ + echo "[make] bridge encerrado (pid $$pid)"; \ + fi; \ + else \ + echo "[make] bridge nao estava rodando, removendo pid stale"; \ + fi; \ + rm -f "$(BRIDGE_PID_FILE)"; \ + else \ + echo "[make] bridge nao esta rodando"; \ + fi + @port_pids="$$(lsof -t -iTCP:$(DEV_WS_PORT) -sTCP:LISTEN 2>/dev/null || true)"; \ + if [ -n "$$port_pids" ]; then \ + for pid in $$port_pids; do \ + kill "$$pid" >/dev/null 2>&1 || true; \ + sleep 1; \ + if kill -0 "$$pid" >/dev/null 2>&1; then \ + kill -9 "$$pid" >/dev/null 2>&1 || true; \ + echo "[make] worker do bridge encerrado com SIGKILL na porta $(DEV_WS_PORT) (pid $$pid)"; \ + else \ + echo "[make] worker do bridge encerrado na porta $(DEV_WS_PORT) (pid $$pid)"; \ + fi; \ + done; \ + fi + +bridge-status: + @if [ -f "$(BRIDGE_PID_FILE)" ] && kill -0 "$$(cat "$(BRIDGE_PID_FILE)")" >/dev/null 2>&1; then \ + echo "bridge: running (pid $$(cat "$(BRIDGE_PID_FILE)"))"; \ + elif [ -n "$$(lsof -t -iTCP:$(DEV_WS_PORT) -sTCP:LISTEN 2>/dev/null || true)" ]; then \ + echo "bridge: running (listener na porta $(DEV_WS_PORT) sem pid file)"; \ + else \ + echo "bridge: stopped"; \ + fi + +bridge-logs: + @if [ ! -f "$(BRIDGE_LOG_FILE)" ]; then \ + echo "ERRO: log do bridge nao encontrado em $(BRIDGE_LOG_FILE)"; \ + exit 1; \ + fi + @tail -f "$(BRIDGE_LOG_FILE)" + +local-up: livekit bridge-up agent-up + @echo "[make] ambiente local iniciado" + @echo "[make] bridge: http://127.0.0.1:$(DEV_WS_PORT)/voice-client" + @echo "[make] livekit: ws://127.0.0.1:$(LIVEKIT_PORT_HTTP)" + @echo "[make] logs: $(RUN_DIR)" + +local-down: + @$(MAKE) --no-print-directory agent-down + @$(MAKE) --no-print-directory bridge-down + @$(MAKE) --no-print-directory livekit-stop + +local-status: + @$(MAKE) --no-print-directory bridge-status + @$(MAKE) --no-print-directory agent-status + @$(MAKE) --no-print-directory livekit-status + +local-logs: + @echo "bridge log: $(BRIDGE_LOG_FILE)" + @echo "agent log: $(AGENT_LOG_FILE)" + +local-stresstest: check-venv + @diag_dump_dir="$${STT_DUMP_DIR:-$${STRESS_REPORT_DIR:-$(RUN_DIR)/local-stresstest}/stt_dumps}"; \ + mkdir -p "$$diag_dump_dir"; \ + if [ "$${STRESS_RESTART_AGENT_FOR_DIAGNOSTICS:-1}" = "1" ]; then \ + $(MAKE) --no-print-directory agent-down; \ + fi; \ + AGENT_DIAG_STT_DUMP_DIR="$$diag_dump_dir" \ + AGENT_DIAG_FLOW_LOG_VAD_DECISIONS="$${FLOW_LOG_VAD_DECISIONS:-1}" \ + AGENT_DIAG_FLOW_LOG_VAD_ACTIVITY="$${FLOW_LOG_VAD_ACTIVITY:-1}" \ + AGENT_DIAG_FLOW_LOG_VAD_ACTIVITY_MIN_PROB="$${FLOW_LOG_VAD_ACTIVITY_MIN_PROB:-0.03}" \ + $(MAKE) --no-print-directory local-up + @set -a; [ -f "$(ENV_DEV)" ] && source "$(ENV_DEV)"; set +a; \ + STT_DUMP_DIR="$${STT_DUMP_DIR:-$${STRESS_REPORT_DIR:-$(RUN_DIR)/local-stresstest}/stt_dumps}" \ + FLOW_LOG_VAD_DECISIONS="$${FLOW_LOG_VAD_DECISIONS:-1}" \ + FLOW_LOG_VAD_ACTIVITY="$${FLOW_LOG_VAD_ACTIVITY:-1}" \ + FLOW_LOG_VAD_ACTIVITY_MIN_PROB="$${FLOW_LOG_VAD_ACTIVITY_MIN_PROB:-0.03}" \ + STRESS_ENV_FILE="$(ENV_DEV)" \ + STRESS_BRIDGE_URL="$${STRESS_BRIDGE_URL:-ws://127.0.0.1:$(DEV_WS_PORT)/ws/agent}" \ + STRESS_BRIDGE_HEALTH_URL="$${STRESS_BRIDGE_HEALTH_URL:-http://127.0.0.1:$(DEV_WS_PORT)/health}" \ + STRESS_AGENT_HEALTH_URL="$${STRESS_AGENT_HEALTH_URL:-http://127.0.0.1:$(AGENT_SERVER_PORT)/}" \ + PYTHONPATH="$(PYTHONPATH_HOST):$$PYTHONPATH" \ + $(PY) -m app.tools.local_stresstest + +# ====================== +# PROD (podman): build + models + up +# ====================== +docker-build: + # build sempre mirando PROD (env file usado no run) e bridge exposto em 8000 + podman build --no-cache -t $(IMAGE) . + +docker-models: docker-build + @podman volume create $(HF_CACHE_VOL) >/dev/null 2>&1 || true + @echo "[make] downloading turn-detector models into volume $(HF_CACHE_VOL)..." + podman run --rm --network host \ + --env-file $(ENV_PROD) \ + -v $(HF_CACHE_VOL):$(HF_HOME_CONT):Z \ + -e HF_HOME=$(HF_HOME_CONT) \ + -e HF_HUB_CACHE=$(HF_HUB_CACHE_CONT) \ + -e PYTHONPATH=$(PYTHONPATH_CONT) \ + $(IMAGE) \ + python -m $(AGENT_FILE) download-files + +docker-up: docker-build docker-models + @podman pod rm -f $(POD_NAME) >/dev/null 2>&1 || true + podman pod create --name $(POD_NAME) --network host + + # AGENT (PROD) - host network + podman run -d --replace --name $(POD_NAME)-agent --pod $(POD_NAME) \ + --env-file $(ENV_PROD) \ + -v $(EXPORT_DIR_HOST):$(EXPORT_DIR_CONT):Z \ + -v $(HF_CACHE_VOL):$(HF_HOME_CONT):Z \ + -e EXPORT_DIR=$(EXPORT_DIR_CONT) \ + -e HF_HOME=$(HF_HOME_CONT) \ + -e HF_HUB_CACHE=$(HF_HUB_CACHE_CONT) \ + -e LIVEKIT_URL=ws://127.0.0.1:$(LIVEKIT_PORT_HTTP) \ + -e PYTHONPATH=$(PYTHONPATH_CONT) \ + $(IMAGE) \ + python -m $(AGENT_FILE) start --log-level info + + # BRIDGE (PROD) - expõe na porta 8000 (host) + podman run -d --replace --name $(POD_NAME)-bridge --pod $(POD_NAME) \ + --env-file $(ENV_PROD) \ + -v $(HF_CACHE_VOL):$(HF_HOME_CONT):Z \ + -e HF_HOME=$(HF_HOME_CONT) \ + -e HF_HUB_CACHE=$(HF_HUB_CACHE_CONT) \ + -e WS_HOST=0.0.0.0 \ + -e WS_PORT=$(PROD_WS_PORT) \ + -e PYTHONPATH=$(PYTHONPATH_CONT) \ + $(IMAGE) \ + python -m $(BRIDGE_MODULE) --host 0.0.0.0 --port $(PROD_WS_PORT) --log-level info + +docker-down: + @podman pod rm -f $(POD_NAME) >/dev/null 2>&1 || true + +docker-status: + @podman pod ps --filter "name=$(POD_NAME)" || true + @podman ps --filter "name=$(POD_NAME)-" || true + +docker-logs: + @podman logs -f $(POD_NAME)-agent + +docker-logs-bridge: + @podman logs -f $(POD_NAME)-bridge + +# ====================== +# LiveKit (host network) +# ====================== +livekit-stop: + @podman rm -f $(LIVEKIT_NAME) >/dev/null 2>&1 || true + +livekit-config-check: + @test -f "$(LIVEKIT_CFG)" || (echo "ERRO: arquivo $(LIVEKIT_CFG) nao existe"; exit 1) + +livekit: livekit-stop livekit-config-check +ifeq ($(OS),Darwin) + podman run -d --name $(LIVEKIT_NAME) \ + -p $(LIVEKIT_PORT_HTTP):7880 \ + -p $(LIVEKIT_PORT_TCP):7881 \ + -v "$(LIVEKIT_CFG):/etc/livekit.yaml:ro" \ + $(LIVEKIT_IMAGE) \ + --config /etc/livekit.yaml + @echo "LiveKit subiu (macOS/Podman: portas publicadas)." + @echo "Signal URL: ws://127.0.0.1:$(LIVEKIT_PORT_HTTP)" + @echo "RTC TCP URL: 127.0.0.1:$(LIVEKIT_PORT_TCP)" +else + podman run -d --name $(LIVEKIT_NAME) \ + --network host \ + -v "$(LIVEKIT_CFG):/etc/livekit.yaml:ro" \ + $(LIVEKIT_IMAGE) \ + --config /etc/livekit.yaml + @echo "LiveKit subiu (host network)." + @echo "Signal URL: ws://127.0.0.1:$(LIVEKIT_PORT_HTTP)" +endif + +livekit-logs: + podman logs -f $(LIVEKIT_NAME) + +livekit-status: + @podman ps --filter "name=$(LIVEKIT_NAME)" + +CONTAINER=api-tia-pod-agent +SHELL_IN=sh + +BACKUP_DIR=backup_logs +DATE=$(shell date +%Y%m%d_%H%M%S) + +copy-logs: + @DEST="backup_logs_$$(date +%Y%m%d_%H%M%S)" && \ + mkdir -p "$$DEST" && \ + podman cp api-tia-pod-agent:/app/logs "$$DEST/logs" || true && \ + podman cp api-tia-pod-agent:/app/timeline "$$DEST/timeline" || true && \ + podman cp api-tia-pod-agent:/app/log_agent "$$DEST/log_agent" || true && \ + echo "Copiado para $$DEST" + diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000..2f46a22 --- /dev/null +++ b/pytest.ini @@ -0,0 +1,16 @@ +[pytest] +minversion = 7.0 +testpaths = tests +python_files = test_*.py +python_classes = Test* +python_functions = test_* +addopts = + --strict-markers + --cov=src + --cov-report=html + --cov-report=term-missing + -v +markers = + unit: Unit tests + integration: Integration tests + property_test: Property-based tests diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..42df424 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,34 @@ +python-dotenv==1.0.1 +fastapi==0.115.6 +uvicorn[standard]==0.34.0 +httpx==0.27.2 +aiohttp>=3.9,<4 +websockets>=12,<16 +google-cloud-pubsub==2.37.0 +opentelemetry-sdk==1.39.1 +opentelemetry-exporter-otlp-proto-http==1.39.1 +redis==8.0.1 + +livekit==1.0.23 +livekit-api==1.1.0 +livekit-agents[azure,openai,elevenlabs,silero,turn-detector]==1.3.10 +livekit-plugins-azure==1.3.10 +livekit-plugins-turn-detector==1.3.10 + +soundfile==0.12.1 + +langgraph==1.0.8 +langchain==1.2.10 +langchain-openai==1.1.10 +langchain-oci==0.2.5 +langfuse==3.10.0 +oci>=2,<3 + +num2words==0.5.14 + +vosk>=0.3.44,<0.4.0 + +elevenlabs==2.22.1 + +pytest>=8,<10 +pytest-cov>=4,<6 diff --git a/scripts/deploy-regional-k8s.sh b/scripts/deploy-regional-k8s.sh new file mode 100644 index 0000000..ec26b12 --- /dev/null +++ b/scripts/deploy-regional-k8s.sh @@ -0,0 +1,15 @@ +#!/usr/bin/env bash +set -euo pipefail +ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +ENV_FILE="${1:-$ROOT/k8s/regional/regions.env}" +OUT="${2:-$ROOT/k8s/regional/rendered}" +"$ROOT/scripts/render-regional-k8s.sh" "$ENV_FILE" "$OUT" +kubectl apply -f "$OUT/service.yaml" +kubectl apply -f "$OUT/pdb.yaml" +kubectl apply -f "$OUT/deployment-ord.yaml" +kubectl apply -f "$OUT/deployment-iad.yaml" +kubectl apply -f "$OUT/hpa-ord.yaml" +kubectl apply -f "$OUT/hpa-iad.yaml" +kubectl rollout status -f "$OUT/deployment-ord.yaml" --timeout=10m +kubectl rollout status -f "$OUT/deployment-iad.yaml" --timeout=10m +kubectl get pods -l app="$(grep '^APP_NAME=' "$ENV_FILE" | cut -d= -f2)-regional" -o wide diff --git a/scripts/render-regional-k8s.sh b/scripts/render-regional-k8s.sh new file mode 100644 index 0000000..8da0992 --- /dev/null +++ b/scripts/render-regional-k8s.sh @@ -0,0 +1,24 @@ +#!/usr/bin/env bash +set -euo pipefail +ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +ENV_FILE="${1:-$ROOT/k8s/regional/regions.env}" +OUT="${2:-$ROOT/k8s/regional/rendered}" +[[ -f "$ENV_FILE" ]] || { echo "Missing $ENV_FILE (copy regions.env.example)" >&2; exit 2; } +set -a; source "$ENV_FILE"; set +a +command -v envsubst >/dev/null || { echo "envsubst is required (gettext package)" >&2; exit 2; } +rm -rf "$OUT"; mkdir -p "$OUT" +render_region() { + local prefix="$1" + export REGION_ID REGION_REPLICAS XAI_UPSTREAM_URL XAI_SECRET_NAME + REGION_ID="$(eval echo \"\${${prefix}_REGION_ID}\")" + REGION_REPLICAS="$(eval echo \"\${${prefix}_REGION_REPLICAS}\")" + XAI_UPSTREAM_URL="$(eval echo \"\${${prefix}_XAI_UPSTREAM_URL}\")" + XAI_SECRET_NAME="$(eval echo \"\${${prefix}_XAI_SECRET_NAME}\")" + envsubst < "$ROOT/k8s/regional/deployment-region.yaml" > "$OUT/deployment-${REGION_ID}.yaml" + envsubst < "$ROOT/k8s/regional/hpa-region.yaml" > "$OUT/hpa-${REGION_ID}.yaml" +} +render_region ORD +render_region IAD +envsubst < "$ROOT/k8s/regional/service.yaml" > "$OUT/service.yaml" +envsubst < "$ROOT/k8s/regional/pdb.yaml" > "$OUT/pdb.yaml" +echo "Rendered manifests in $OUT" diff --git a/scripts/validate-regional-k8s.sh b/scripts/validate-regional-k8s.sh new file mode 100644 index 0000000..9c70e92 --- /dev/null +++ b/scripts/validate-regional-k8s.sh @@ -0,0 +1,11 @@ +#!/usr/bin/env bash +set -euo pipefail +ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +ENV_FILE="${1:-$ROOT/k8s/regional/regions.env}" +OUT="${2:-$ROOT/k8s/regional/rendered}" +"$ROOT/scripts/render-regional-k8s.sh" "$ENV_FILE" "$OUT" +for f in "$OUT"/*.yaml; do + echo "== validating $(basename "$f") ==" + kubectl apply --dry-run=client -f "$f" >/dev/null + echo OK +done diff --git a/sh/run_export.sh b/sh/run_export.sh new file mode 100644 index 0000000..0149245 --- /dev/null +++ b/sh/run_export.sh @@ -0,0 +1,13 @@ +#!/bin/bash + +ORACLE_HOME=/opt/oracle/instantclient_21_9 +export PATH=$ORACLE_HOME:$PATH LD_LIBRARY_PATH=$ORACLE_HOME + +sqlplus -s admin/@ @export_data_pump.sql + +if [ $? -ne 0 ]; then + echo "Export falhou em $(date)" + exit 1 +else + echo "Export concluído com sucesso em $(date)" +fi \ No newline at end of file diff --git a/src/agent/base/base_classifier.py b/src/agent/base/base_classifier.py new file mode 100644 index 0000000..2fee1fd --- /dev/null +++ b/src/agent/base/base_classifier.py @@ -0,0 +1,66 @@ +from __future__ import annotations + +import warnings +from functools import lru_cache +from importlib import resources +from typing import Any, Dict + +from langchain_core.prompts import PromptTemplate +from langchain_openai import ChatOpenAI +from langchain_oci import ChatOCIGenAI +from langchain_core.runnables import RunnablePassthrough + +warnings.filterwarnings("ignore") + +from dotenv import load_dotenv +import os +load_dotenv() + +class BaseClassifier: + MODEL_NAME = "openai/gpt-oss-20b" + API_KEY = "fake-key" + BASE_URL = "http://10.153.34.154/gpt-oss-20b/v1" + REASONING_EFFORT = "low" + PROMPT_FILE: str = "" + + def __init__(self, *, temperature: float = 0.0) -> None: + self.llm = ChatOCIGenAI( + model_id=os.getenv("OCI_ENDPOINT_ID", ""), + service_endpoint=os.getenv("OCI_ENDPOINT", ""), + compartment_id=os.getenv("OCI_COMPARTMENT_ID", ""), + auth_file_location=os.getenv("OCI_AUTH_FILE_LOCATION", "./config"), + model_kwargs={"temperature": 0.0, + "top_p": 0.1, + "reasoning_effort":"MINIMAL"} + ) + """self.llm = ChatOpenAI( + model="openai/gpt-oss-20b", + api_key="fake-key", + base_url="http://10.153.34.154/gpt-oss-20b/v1", + temperature=0.0, + top_p=0.1, + reasoning_effort="low" + )""" + + self.prompt = PromptTemplate( + template=self._load_prompt(), + input_variables=["text"], + partial_variables=self._partial_variables(), + ) + + self.chain = ( + {"text": RunnablePassthrough()} + | self.prompt + | self.llm + ) + + @classmethod + @lru_cache + def _load_prompt(cls) -> str: + return resources.files("agent.prompts").joinpath(cls.PROMPT_FILE).read_text() + + def _partial_variables(self) -> Dict[str, Any]: + return {} + + def run(self, text: str) -> Dict[str, Any]: + raise NotImplementedError("Subclasses devem implementar o método run()") diff --git a/src/agent/base/base_stage.py b/src/agent/base/base_stage.py new file mode 100644 index 0000000..79e7758 --- /dev/null +++ b/src/agent/base/base_stage.py @@ -0,0 +1,205 @@ +# base/base.py +import os +import time +from langgraph.prebuilt import create_react_agent +from langchain_core.prompts import ChatPromptTemplate +from langchain_core.messages import HumanMessage, AIMessage, SystemMessage, trim_messages +from langchain_core.callbacks import BaseCallbackHandler +from langchain_openai import ChatOpenAI +from langchain_oci import ChatOCIGenAI +#from langfuse import get_client +#from langfuse.langchain import CallbackHandler as LangfuseCallbackHandler +from functools import lru_cache +import queue, threading + +from dotenv import load_dotenv +import os +load_dotenv() + +class StreamCaptureHandler(BaseCallbackHandler): + def __init__(self): + self.tokens = [] + self.buffer = "" + self.final_result = None + self.steps = [] + + def on_llm_new_token(self, token, **kwargs): + self.tokens.append(token) + + def on_chain_end(self, outputs, **kwargs): + self.final_result = outputs + + def on_tool_end(self, output, **kwargs): + self.steps.append(output) + +@lru_cache(maxsize=32) +def load_langfuse_prompt( + prompt_name: str, + ttl_seconds: int = 300, +) -> str: + """ + Carrega prompt do Langfuse com cache por processo + e revalidação automática a cada 5 minutos. + """ + prompt = langfuse.get_prompt( + name=prompt_name, + cache_ttl_seconds=ttl_seconds, + ) + return prompt.prompt + +#langfuse = get_client() +#lf_handler = LangfuseCallbackHandler() + +class BaseAgent: + def __init__(self, + tools: list, + streaming: bool = False, + prompt_vars: dict | None = None, + agent_name: str = None, + dynamic_prompt: bool = False): + + self.streaming = streaming + self.prompt_vars = prompt_vars or {} + self.agent_name = agent_name + self.dynamic_prompt = dynamic_prompt + + self.llm = ChatOCIGenAI( + model_id=os.getenv("OCI_ENDPOINT_ID", ""), + service_endpoint=os.getenv("OCI_ENDPOINT", ""), + compartment_id=os.getenv("OCI_COMPARTMENT_ID", ""), + auth_file_location=os.getenv("OCI_AUTH_FILE_LOCATION", "./config"), + model_kwargs={"temperature": 0.0, + "top_p": 0.1, + "reasoning_effort":"MINIMAL"} + ) + """self.llm = ChatOpenAI( + model="openai/gpt-oss-20b", + api_key="fake-key", + base_url="http://10.153.34.154/gpt-oss-20b/v1", + temperature=0.0, + top_p=0.1, + reasoning_effort="low", + streaming=self.streaming, + )""" + + self.stream_handler = StreamCaptureHandler() + + raw_prompt = self._load_prompt() + #raw_prompt = load_langfuse_prompt(self.agent_name) + + # Aplica variáveis parciais no template do prompt + if self.prompt_vars: + for key, value in self.prompt_vars.items(): + raw_prompt = raw_prompt.replace("{" + key + "}", str(value)) + + self.system_prompt = raw_prompt + + # Histórico de mensagens gerenciado manualmente (substitui ConversationBufferWindowMemory) + self.messages: list = [] + self.k = 30 # window size — últimas k*2 mensagens (k pares human/ai) + + # Cria o agente usando langgraph-prebuilt create_react_agent + # Se dynamic_prompt=True, usa lambda que lê self.system_prompt a cada invocação + # permitindo que o prompt seja atualizado entre chamadas (ex: UnifiedAgent) + if self.dynamic_prompt: + self.agent_exec = create_react_agent( + model=self.llm, + tools=tools, + prompt=lambda state: [SystemMessage(content=self.system_prompt)] + state["messages"], + ) + else: + self.agent_exec = create_react_agent( + model=self.llm, + tools=tools, + prompt=self.system_prompt, + ) + + def _load_prompt(self) -> str: + raise NotImplementedError + + def _get_trimmed_messages(self) -> list: + """Retorna as últimas k*2 mensagens (janela deslizante).""" + max_messages = self.k * 2 + if len(self.messages) > max_messages: + return self.messages[-max_messages:] + return list(self.messages) + + def inject_user_message(self, text: str): + self.messages.append(HumanMessage(content=text)) + + def inject_ai_message(self, text: str): + self.messages.append(AIMessage(content=text)) + + # ================= RUN NORMAL ====================== + def run(self, user_input: str): + # Adiciona a mensagem do usuário ao histórico + self.messages.append(HumanMessage(content=user_input)) + + # Prepara mensagens com window trimming + input_messages = self._get_trimmed_messages() + + result = self.agent_exec.invoke( + {"messages": input_messages}, + #config={"callbacks": [self.stream_handler, lf_handler]} + ) + + #print("Texto:",result["messages"][-1].content) + #print("Metadata:",result["messages"][-1].usage_metadata) + #print("*"*100) + + # Extrai a última mensagem do AI do resultado + output_messages = result.get("messages", []) + ai_response = "" + tool_steps = [] + + for msg in output_messages: + if isinstance(msg, AIMessage) and msg.content and not getattr(msg, 'tool_calls', None): + ai_response = msg.content + + # Coleta intermediate tool steps + for msg in output_messages: + if hasattr(msg, 'tool_calls') and msg.tool_calls: + tool_steps.extend(msg.tool_calls) + + # Adiciona a resposta do AI ao histórico + if ai_response: + ai_response = ai_response.split('###commentary')[0].strip() # remove commentary + self.messages.append(AIMessage(content=ai_response)) + + #print(self.messages) + + return { + "output": ai_response, + "tools": tool_steps + } + + # ================= STREAMING ====================== + def stream_run(self, user_input: str): + self.stream_handler.tokens = [] + + # Adiciona a mensagem do usuário ao histórico + self.messages.append(HumanMessage(content=user_input)) + + # Prepara mensagens com window trimming + input_messages = self._get_trimmed_messages() + + # Usa streaming do agente + ai_response = "" + for chunk in self.agent_exec.stream( + {"messages": input_messages}, + config={"callbacks": [self.stream_handler]}#, lf_handler + ): + # Processa chunks de streaming + for node_name, node_output in chunk.items(): + if node_name == "agent" and "messages" in node_output: + for msg in node_output["messages"]: + if isinstance(msg, AIMessage) and msg.content: + ai_response = msg.content + for token in msg.content: + yield token + + # Adiciona a resposta ao histórico + if ai_response: + self.messages.append(AIMessage(content=ai_response)) + + yield "" diff --git a/src/agent/classifier/check_guardrail.py b/src/agent/classifier/check_guardrail.py new file mode 100644 index 0000000..33adddd --- /dev/null +++ b/src/agent/classifier/check_guardrail.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +import json +import re +from typing import Any, Dict + +from langchain_core.output_parsers import PydanticOutputParser +from pydantic import BaseModel + +from agent.base.base_classifier import BaseClassifier + + +def get_json(text: str) -> str: + pattern = r"```json(.*?)```" + match = re.search(pattern, text, re.DOTALL) + if match: + return match.group(1).strip() + raise ValueError("JSON não encontrado na resposta do modelo.") + + +class GuardrailResult(BaseModel): + is_attack: bool + + +parser = PydanticOutputParser(pydantic_object=GuardrailResult) + + +class CheckGuardrail(BaseClassifier): + PROMPT_FILE = "check_guardrail.txt" + REASONING_EFFORT = "low" + + def _partial_variables(self) -> Dict[str, Any]: + return { + "format_instructions": parser.get_format_instructions() + } + + def run(self, text: str) -> Dict[str, Any]: + output = self.chain.invoke(text) + json_text = get_json(output.content) + return parser.parse(json_text).dict() diff --git a/src/agent/classifier/conversation_classifier.py b/src/agent/classifier/conversation_classifier.py new file mode 100644 index 0000000..da2bd90 --- /dev/null +++ b/src/agent/classifier/conversation_classifier.py @@ -0,0 +1,80 @@ +from __future__ import annotations + +import json +from datetime import datetime +from typing import Any, Dict + +import langfuse + +from agent.base.base_classifier import BaseClassifier + + +client = langfuse.Langfuse(timeout=30) + + +class ConversationClassifier(BaseClassifier): + PROMPT_FILE = "conversation_classifier.txt" + REASONING_EFFORT = "low" + + def run(self, text: str) -> Dict[str, Any]: + output = self.chain.invoke(text) + return json.loads(output.content) + + +def get_historic_conversation(trace_id): + trace = client.api.trace.get(trace_id).dict() + ordered_obs = sorted( + trace["observations"], + key=lambda x: x.get("startTime") or "" + ) + + conversation = "" + for obs in ordered_obs: + meta = obs.get("metadata", {}) + stage = meta.get("stage") + + if not stage or stage == "check_guardrail": + continue + + model_msg = obs.get("output") + human_msg = obs.get("input") + + if isinstance(human_msg, dict): + human_msg = human_msg.get("text") + + conversation += ( + f"ESTÁGIO:{stage}\n" + f"HUMANO:{human_msg}\n" + f"MODELO:{model_msg}\n\n" + ) + + return conversation.strip() + + +def format_success_purchase(raw_mailing): + allowed_keys = [ + 'NUM_TELEFONE', 'DATA_VENCIMENTO', 'NOME_CLIENTE_COMPLETO', + 'NUM_CPF_CNPJ_CLIENTE', 'ACAO', 'BONUS_DESTINO', 'FLG_FID', + 'USO_DE_DADOS', 'DESCONTO_DESTINO', 'DAT_NASCIMENTO', + 'VALOR_PLANO_CORE', 'VLR_PLANO_DESTINO', + 'VLR_FINAL_PLANO_DEST', 'DADOS_CORE', 'ENDERECO', + 'UF', 'TIPO_MAILING', 'AGING_PLANO' + ] + return {k: v for k, v in raw_mailing.items() if k in allowed_keys} + + +def format_call(raw_mailing, reason_status, summary): + return { + "DATA": datetime.now().strftime("%Y-%m-%d"), + "HORA": datetime.now().strftime("%H:%M:%S"), + "NUM_ACESSO": raw_mailing["NUM_TELEFONE"], + "NOME": raw_mailing["NOME_CLIENTE_COMPLETO"], + "NUM_DOCUMENTO": raw_mailing["NUM_CPF_CNPJ_CLIENTE"], + "STATUS": "CONTATO", + "DESC_STATUS": list(reason_status.keys())[0], + "MOTIVO_STATUS": list(reason_status.values())[0], + "RESUMO": summary, + "QTD_TENTATIVAS": 1, + "NOME_MAILING": raw_mailing["MAILING"], + "SEGMENTACAO1": "CONTROLE_POS", + } diff --git a/src/agent/classifier/conversation_summary.py b/src/agent/classifier/conversation_summary.py new file mode 100644 index 0000000..613a421 --- /dev/null +++ b/src/agent/classifier/conversation_summary.py @@ -0,0 +1,13 @@ +from __future__ import annotations + +from typing import Any, Dict + +from agent.base.base_classifier import BaseClassifier + +class ConversationSummary(BaseClassifier): + PROMPT_FILE = "conversation_summary.txt" + REASONING_EFFORT = "low" + + def run(self, text: str) -> str: + output = self.chain.invoke(text) + return output.content diff --git a/src/agent/classifier/interruption_classifier.py b/src/agent/classifier/interruption_classifier.py new file mode 100644 index 0000000..398522f --- /dev/null +++ b/src/agent/classifier/interruption_classifier.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +import json +import re +from typing import Any, Dict + +from langchain_core.output_parsers import PydanticOutputParser +from pydantic import BaseModel + +from agent.base.base_classifier import BaseClassifier +from app.utils.logging import setup_minimal_logging +from app.common.timed import timed +logger = setup_minimal_logging() + +class InterruptionClassifier(BaseClassifier): + PROMPT_FILE = "interruption.txt" + REASONING_EFFORT = "low" + + @timed("InterruptionClassifier.run") + def run(self, text: str) -> bool: + output = self.chain.invoke(text) + logger.info(f"Resultado modelo interrupção: {text}, {bool(eval(output.content))}") + return bool(eval(output.content)) \ No newline at end of file diff --git a/src/agent/pipeline/__main__.py b/src/agent/pipeline/__main__.py new file mode 100644 index 0000000..a97169c --- /dev/null +++ b/src/agent/pipeline/__main__.py @@ -0,0 +1,164 @@ +from pipeline.customer_pipeline import CustomerPipeline + +""" +EXECUÇÃO INTERATIVA VIA TERMINAL +> python -m pipeline +""" + + +def load_mock_data(): + """ + Você pode substituir isso por integração real com seu backend. + Aqui fica só um exemplo de estrutura esperada. + """ + mailing = { + "NUM_TELEFONE": 99999999999, + "DATA_VENCIMENTO": "VENC DIA 20", + "NOME_CLIENTE_COMPLETO": "Maria de Fátima da silva", + "NUM_CPF_CNPJ_CLIENTE": "00000000123", + "NOME_MAE": "Mariana da silva", + "VLR_TM_ORIGEM": 79.99, + "PLANO_ORIGEM": " TIM CONTROLE SMART 8 0", + "PLANO_DESTINO": "TIM_BLACK_A", + "ACAO": "CTRL SMART NAO FIDEL_POS ATL A FIDEL", + "BONUS_DESTINO": "15+30", + "PARCEIRO": "ATN", + "FLG_FID": 0, + "USO_DE_DADOS": "Até 500MB", + "DESCONTO_DESTINO": "90 FIDEL", + "PRIORIZACAO": 1, + "CUSTCODE": 1314082913.0, + "TP_BASE": "ESTRUTURAL", + "FLAG_DIRECIONAMENTO": "PADRAO", + "FRESCOR": "M-2", + "MEDIA_RECARGA_2": 67.99, + "DAT_NASCIMENTO": "1968-09-08", + "FAIXA_RECARGA": "61 |- 70", + "FLG_LIVE_MKT": None, + "BONUS_CONVERGENCIA": 0, + "ORDEM": 4.0, + "VALOR_PLANO_CORE": 92.99, + "APP_STR": None, + "OFERTA_LIVE": None, + "IDADE": 56.0, + "DADOS_PERIODO": "TARDE", + "CLASSIFICACAO_RISCO": "BAIXO RISCO", + "VLR_PLANO_DESTINO": 189.99, + "FLG_ICMS": 0, + "DESCONTO_ICMS": 0, + "VLR_FINAL_PLANO_DEST": 99.99, + "FX_IDADE": "51 A 60 ANOS", + "PERSONA_IDADE": "ADULTO EXPERIENTE", + "FX_GB": "02. ATÉ 1GB", + "PERSONA_GB": "0-1 USO MINIMO", + "PERSONA_USO": "TARDES CONECTADAS", + "PERSONA_UF": "METROPOLE DINÂMICA", + "PERSONA_SOCIAL_MEDIA": 0, + "PESONA_STREAMER": 0, + "PERSONA_MUSIC": 0, + "DADOS_CORE": "5GB", + "FX_PROPENSAO": 17.0, + "ENDERECO": "N/A SQS 115 BL I 0 - CEP: 70385090, Cidade: BRASILIA/DF", + "UF": "DF", + "FLG_TRIPLE_A_OP": 0, + "TIPO_MAILING": "CTRL-POS", + "ANOMES": 202509, + "CANAL": "OUTBOUND", + "AGING_PLANO": 39.0, + "DELTA_TICKET": 32.0 + } + actual_plan = { + "plano": "Tim Controle Smart oito ponto zero", + "beneficios_001": None, + "beneficios_002": None, + "beneficios_todos": None, + "apps_zero_rating": "WhatsApp, Messenger, Instagram, Facebook, X", + "categoria_plano": "controle", + "dependentes": None, + "grupo_plano": "Tim Controle (Fatura)", + "forma_pagamento": "Fatura", + "dados_GB": "5", + "comunicacao_whatsapp": None, + "comunicacao_ligacoes_voz": "Ilimitadas", + "comunicacao_sms": "Ilimitados", + "inflight": None, + "roaming_internacional": None, + "servicos_valor_agregado": "Aya Books Premium, Aya Ensinah Premium, EXA Segurança", + "aplicativos_inclusos": None + } + target_plan = { + "plano": "Tim Black C Light", + "beneficios_001": "não consumir seus gigas quando usar os aplicativos Instagram, Facebook e X e quando trocar mensagens no whatsapp.", + "beneficios_002": "não consome seus gigas quando usar as principais redes sociais e tem acesso ao wifi nos aviões da gol e da latam e tem também o pacote chile de roaming internacional.", + "beneficios_todos": "Utilização de ligações SMS, WhatsApp, Instagram, Facebook e X ilimitados, Além acesso ao wifi nos aviões da gol e latam juntamente com o pacote chile de roaming internacional. Além disso diversos serviços extras como Aya Audiobooks Premium, Bancah Premium + Jornais, EXA Segurança Premium, Aya Ensinah Premium, EXA Cloud, Busuu, Fluid Premium", + "apps_zero_rating": "WhatsApp, Instagram, Facebook, X", + "categoria_plano": "pós pago", + "dependentes": "Não há dependentes", + "grupo_plano": "Tim Black", + "forma_pagamento": "Fatura", + "dados_GB": "20", + "comunicacao_whatsapp": "O uso de audio e video consome dos seus gigas, enquanto mensagens são ilimitadas", + "comunicacao_ligacoes_voz": "Ilimitadas", + "comunicacao_sms": "Ilimitados", + "inflight": "Tim no Avião", + "roaming_internacional": "Pacote Chile", + "servicos_valor_agregado": "Aya Audiobooks Premium, Bancah Premium + Jornais, EXA Segurança Premium, Aya Ensinah Premium, EXA Cloud, Busuu, Fluid Premium", + "aplicativos_inclusos": "Não há aplicativo incluso" + } + + return mailing, actual_plan, target_plan + + +def main(): + print("\n==============================") + print(" CUSTOMER PIPELINE (CLI) ") + print("==============================\n") + + # ---------------------------------------------- + # Carrega dados fake ou reais + # ---------------------------------------------- + mailing, actual_plan, target_plan = load_mock_data() + + # ---------------------------------------------- + # Cria o pipeline + # ---------------------------------------------- + pipeline = CustomerPipeline( + mailing=mailing, + actual_plan=actual_plan, + target_plan=target_plan, + streaming=True + ) + + # ---------------------------------------------- + # Prepara (elegibilidade vem do backend real) + # ---------------------------------------------- + intro = pipeline.start() + pipeline.prepare(elegibility=True, protocol = "PRT-2025-0000123456") + # ---------------------------------------------- + # Inicia com mensagem fake do LLM + # ---------------------------------------------- + print(f"[{intro['stage'].upper()}] AGENTE:", intro["response"]) + + # ---------------------------------------------- + # LOOP PRINCIPAL (CLI) + # ---------------------------------------------- + while True: + user = input("\nVocê: ") + + if user.lower() in ["sair", "exit", "quit"]: + print("\nEncerrando manualmente.") + break + + result = pipeline.run(user) + + print(f"[{result['stage'].upper()}] AGENTE:", result["response"]) + print("TOOLS :", result["tools"]) + print("STATE :", result["state"]) + + if result["stage"] == "end": + print("\nPipeline concluído.\n") + break + + +if __name__ == "__main__": + main() diff --git a/src/agent/pipeline/customer_pipeline_langgraph.py b/src/agent/pipeline/customer_pipeline_langgraph.py new file mode 100644 index 0000000..9e5c39b --- /dev/null +++ b/src/agent/pipeline/customer_pipeline_langgraph.py @@ -0,0 +1,942 @@ +from typing import TypedDict, Any, Optional, Tuple +from langgraph.graph import StateGraph, END + +from langfuse import get_client +from agent.stage.argumentation import ArgumentationAgent +from agent.stage.data_confirmation import DataConfirmationAgent +from agent.stage.formalization import FormalizationAgent +from agent.stage.presentation import PresentationAgent +from agent.utils.utils import ( + process_mailing, + process_plan, + get_gender, + get_plan, + normalize_name, + sentence_confidence, + start_message_argumentation, +) + +from agent.classifier.check_guardrail import CheckGuardrail +from agent.classifier.conversation_summary import ConversationSummary +from agent.classifier.conversation_classifier import ( + get_historic_conversation, + ConversationClassifier, + format_success_purchase, + format_call, +) + +from app.common.timed import timed +from langfuse.langchain import CallbackHandler +import asyncio +import copy +import os +import json +import math +import time +from datetime import datetime +from pathlib import Path +from app.utils.logging import setup_minimal_logging + +logger = setup_minimal_logging() +langfuse = get_client() +langfuse_handler = CallbackHandler() + +# ===================== CONTROLE DE LOG DO AGENTE ===================== +ENABLE_AGENT_LOG = True +# ===================================================================== + + +class AgentLogger: + """Logger que grava em TXT todo o ciclo de vida do pipeline do agente.""" + + def __init__(self, telefone: str): + self.telefone = str(telefone) + self.start_ts = datetime.now() + ts_str = self.start_ts.strftime("%H%M%S") + + base_dir = Path("log_agent") / self.start_ts.strftime("%Y-%m-%d") / self.start_ts.strftime("%H") + base_dir.mkdir(parents=True, exist_ok=True) + + self.filepath = base_dir / f"{ts_str}_{self.telefone}.txt" + self._lines: list[str] = [] + self._write_header() + + # ---- helpers internos ---- + def _write_header(self): + self._append("=" * 80) + self._append(f"AGENT LOG — Telefone: {self.telefone}") + self._append(f"Início da sessão: {self.start_ts.isoformat()}") + self._append("=" * 80) + self._append("") + + def _ts(self) -> str: + return datetime.now().strftime("%H:%M:%S.%f")[:-3] + + def _append(self, text: str): + self._lines.append(text) + + def _flush(self): + with open(self.filepath, "w", encoding="utf-8") as f: + f.write("\n".join(self._lines)) + + def _format_messages(self, messages: list) -> str: + parts = [] + for msg in messages: + role = type(msg).__name__.replace("Message", "").upper() + content = getattr(msg, "content", str(msg)) + tool_calls = getattr(msg, "tool_calls", None) + line = f" [{role}] {content}" + if tool_calls: + line += f" | tool_calls={tool_calls}" + parts.append(line) + return "\n".join(parts) if parts else " (vazio)" + + def _format_state(self, state_obj) -> str: + if state_obj is None: + return " (sem state)" + d = {k: v for k, v in state_obj.__dict__.items() if not k.startswith("_")} + parts = [f" {k}: {v}" for k, v in d.items()] + return "\n".join(parts) + + def _format_dict(self, d: dict, indent: int = 4) -> str: + prefix = " " * indent + parts = [] + for k, v in d.items(): + parts.append(f"{prefix}{k}: {v}") + return "\n".join(parts) if parts else f"{prefix}(vazio)" + + # ---- métodos públicos de logging ---- + + def log_section(self, title: str): + self._append("") + self._append("-" * 80) + self._append(f"[{self._ts()}] {title}") + self._append("-" * 80) + self._flush() + + def log(self, label: str, content: str = ""): + if content: + self._append(f" [{self._ts()}] {label}: {content}") + else: + self._append(f" [{self._ts()}] {label}") + self._flush() + + def log_system_prompt(self, agent_name: str, prompt: str): + self._append(f" [{self._ts()}] SYSTEM PROMPT ({agent_name}):") + for line in prompt.splitlines(): + self._append(f" | {line}") + self._flush() + + def log_messages(self, label: str, messages: list): + self._append(f" [{self._ts()}] {label} — Mensagens ({len(messages)}):") + self._append(self._format_messages(messages)) + self._flush() + + def log_agent_state(self, stage: str, state_obj): + self._append(f" [{self._ts()}] STATE ({stage}):") + self._append(self._format_state(state_obj)) + self._flush() + + def log_metadata(self, metadata: dict): + self._append(f" [{self._ts()}] METADATA:") + self._append(self._format_dict(metadata)) + self._flush() + + def log_stage_result(self, stage: str, output: str, next_stage: str, auto: bool): + self._append(f" [{self._ts()}] RESULTADO do estágio '{stage}':") + self._append(f" output: {output}") + self._append(f" próximo estágio: {next_stage}") + self._append(f" auto: {auto}") + self._flush() + + +class PipelineState(TypedDict, total=False): + stage: str + user_input: dict + auto: bool + output: str + is_attack: bool + + +class CustomerPipeline: + def __init__(self, mailing, streaming=False): + self.streaming = streaming + self.stage = "presentation" + self.classified_conversation = None + self.accepted_purchase = False + self.name = None + self.attacks = 0 + + self.prompt_vars = { + "raw_customer": mailing, + "customer": None, + "actual_plan": None, + "target_plan": None, + "elegibility": None, + } + + self.trace_id = langfuse.create_trace_id() + self.intro = None + self.graph = None + self.agent = None + + # Buffer do último output (aguarda sinal do backend) + self.pending_input = None + self.pending_output = None + self.pending_stage = None + self.pending_metadata = None + + # Logger do agente (TXT) + telefone = mailing.get("NUM_TELEFONE", "sem_telefone") + self.agent_log: AgentLogger | None = AgentLogger(telefone) if ENABLE_AGENT_LOG else None + + # ==================================================================================== + # Buffer + commit no Langfuse + + def _buffer_langfuse_trace(self, input_value: Any, output_value: str): + """ + Guarda o último par (input, output) gerado pelo agente para só subir no Langfuse + quando o backend sinalizar se foi ouvido por completo ou interrompido. + """ + self.pending_input = input_value + self.pending_output = output_value + self.pending_stage = self.stage + self.pending_metadata = self.get_formatted_metadata() + + def update_langfuse_trace(self, final_output: str, was_interrupted: bool = False): + """ + Faz o commit no Langfuse usando os campos pendentes (input/stage/metadata) + e o output final (completo ou truncado). + """ + #print('-'*30) + #print(final_output,was_interrupted) + #print(self.pending_input) + if self.pending_stage is None: + return + + metadata = self.pending_metadata or {} + metadata = {**metadata, "interrupted": was_interrupted} + + with langfuse.start_as_current_observation( + as_type="span", + name=self.pending_stage, + trace_context={"trace_id": self.trace_id}, + ) as span: + langfuse.update_current_span( + metadata=metadata, + name=self.pending_stage, + output=final_output, + input=self.pending_input, + ) + + def update_langfuse_auto(self, input, output): + with langfuse.start_as_current_observation( + as_type="span", + name=self.stage, + trace_context={"trace_id": self.trace_id}, + ) as span: + langfuse.update_current_span( + metadata=self.get_formatted_metadata(), + name=self.stage, + output=output, + input=input + ) + + # ==================================================================================== + # Backend sinaliza consumo/interrupção + + def set_interruption( + self, + was_interrupted: bool, + listened_text: Optional[str] = None, + skipped_vacalization: bool = False): + + logger.info( + "SET_INTERRUPTION | was_interrupted=%s | listened_text=%r | skipped_vacalization=%s", + was_interrupted, + listened_text, + skipped_vacalization, + ) + + if not self.pending_output: + return + + # Definimos a marca de interrupção + INTERRUPTION_TAG = "###interrupção###" + + if skipped_vacalization: + # Caso 1: O áudio nem chegou a tocar (interrupção imediata ou erro) + final_output = INTERRUPTION_TAG + was_interrupted_flag = True + self.agent.state.restore_state() + self.agent.set_interruption_message(final_output) + + elif was_interrupted: + # Caso 2: Estava falando e foi cortado + # Se ouviu algo, usa o trecho + tag. Se não ouviu nada, apenas a tag. + final_output = f"{listened_text or ''} {INTERRUPTION_TAG}".strip() + was_interrupted_flag = True + self.agent.set_interruption_message(final_output) + + else: + # Caso 3: Fluxo normal (não foi interrompido) + final_output = self.pending_output + was_interrupted_flag = False + + # Log de interrupção + if self.agent_log: + self.agent_log.log_section("INTERRUPÇÃO") + self.agent_log.log("was_interrupted", str(was_interrupted)) + self.agent_log.log("skipped_vacalization", str(skipped_vacalization)) + self.agent_log.log("listened_text", str(listened_text)) + self.agent_log.log("final_output", final_output) + self.agent_log.log("was_interrupted_flag", str(was_interrupted_flag)) + self.agent_log.log_agent_state(self.stage, self.agent.state if self.agent else None) + + # Atualiza o Rastreamento (Langfuse) + self.update_langfuse_trace( + final_output, + was_interrupted=was_interrupted_flag, + ) + + # Limpa buffer + self.pending_input = None + self.pending_output = None + self.pending_stage = None + self.pending_metadata = None + + + # ==================================================================================== + def start(self): + first_name = self.prompt_vars["raw_customer"]["NOME_CLIENTE_COMPLETO"].split()[0].lower() + self.intro = ( + f"Olá, meu nome é Helena, sou consultora de vendas da TIM, para sua segurança a ligação está sendo gravada. " + f"Eu estou falando com {first_name}??" + ) + self.stage = "presentation" + self.name = "presentation" + + if self.agent_log: + self.agent_log.log_section("START") + self.agent_log.log("stage", self.stage) + self.agent_log.log("intro", self.intro) + + return self.stage, self.intro + + def argumentation_start(self): + + text = start_message_argumentation({ + "cliente_alvo_primeiro_nome": self.prompt_vars["customer"]["cliente_alvo_primeiro_nome"].lower(), + "meses_restantes_fidelizacao": self.prompt_vars["actual_plan"]["meses_restantes_fidelizacao"], + "quanto_pago_a_mais": self.prompt_vars["target_plan"]["quanto_pago_a_mais"], + + "quanto_pago_a_mais_int": abs(math.ceil(self.prompt_vars["raw_customer"]["VLR_FINAL_PLANO_DEST"] - self.prompt_vars["raw_customer"]["MEDIA_RECARGA_2"])), + + "gb_plano_atual": self.prompt_vars["actual_plan"]["gb_plano_atual"], + "valor_plano_sem_fidelização": self.prompt_vars["actual_plan"]["valor_plano_sem_fidelização"], + "dados_GB": self.prompt_vars["target_plan"]["dados_GB"], + "beneficios_001": self.prompt_vars["target_plan"]["beneficios_001"], + "valor_plano_final": self.prompt_vars["target_plan"]["valor_plano_final"], + "gb_alvo_diferenca": self.prompt_vars["target_plan"]["gb_alvo_diferenca"], + "preco_por_dia": self.prompt_vars["target_plan"]["preço_por_dia"] + }) + self.agent.inject_user_message("###START###") + self.agent.inject_ai_message(text) + return text + + # ==================================================================================== + def prepare(self, elegibility: bool, protocol: str): + self.stage = "presentation" + self.name = "presentation" + + self.prompt_vars["customer"] = copy.deepcopy(self.prompt_vars["raw_customer"]) + self.prompt_vars["customer"]["PROTOCOLO"] = protocol + + self.prompt_vars["elegibility"] = elegibility + + self.prompt_vars["actual_plan"] = process_plan({}, self.prompt_vars["customer"], "actual") + self.prompt_vars["target_plan"] = process_plan( + get_plan(self.prompt_vars["customer"]["PLANO_DESTINO"]), + self.prompt_vars["customer"], + "target", + ) + + self.prompt_vars["customer"] = get_gender(self.prompt_vars["customer"]) + self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"] = normalize_name( + self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"] + ) + + self.prompt_vars["customer"] = process_mailing(self.prompt_vars["customer"]) + + self.agent = PresentationAgent( + streaming=self.streaming, + prompt_vars={ + "customer_first_name": self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"] + .split()[0] + .lower(), + "elegibility": self.prompt_vars["elegibility"], + }, + ) + + self.agent.inject_user_message("###START###") + self.agent.inject_ai_message(self.intro) + + self.check_guardrail = CheckGuardrail() + self.conversation_classifier = ConversationClassifier() + self.conversation_summary = ConversationSummary() + + self.update_langfuse_auto({"text": "###START###", "prompt_vars": self.prompt_vars}, + self.intro,) + + # Log do prepare + if self.agent_log: + self.agent_log.log_section("PREPARE") + self.agent_log.log("elegibility", str(elegibility)) + self.agent_log.log("protocol", protocol) + self.agent_log.log_system_prompt("presentation", self.agent.system_prompt) + self.agent_log.log_messages("Mensagens iniciais (presentation)", self.agent.messages) + self.agent_log.log("prompt_vars", str(self.prompt_vars)) + + self.graph = self._build_graph() + + # ==================================================================================== + def _build_graph(self): + g = StateGraph(PipelineState) + + g.add_node("guardrail", self._lg_guardrail) + g.add_node("presentation", self._lg_presentation) + g.add_node("argumentation", self._lg_argumentation) + g.add_node("data_confirmation", self._lg_data_confirmation) + g.add_node("formalization", self._lg_formalization) + + g.set_entry_point("guardrail") + + g.add_conditional_edges( + "guardrail", + self._lg_route_from_guardrail, + { + "presentation": "presentation", + "argumentation": "argumentation", + "data_confirmation": "data_confirmation", + "formalization": "formalization", + "end": END, + }, + ) + + for node in ["presentation", "argumentation", "data_confirmation", "formalization"]: + g.add_conditional_edges( + node, + self._lg_route_after_stage, + { + "presentation": "presentation", + "argumentation": "argumentation", + "data_confirmation": "data_confirmation", + "formalization": "formalization", + "end": END, + }, + ) + + return g.compile() + + # ==================================================================================== + # STT normalization + + def _normalize_stt_input(self, user_input: Any) -> Tuple[str, Optional[list]]: + """ + api_text -> retorna (texto, None) + raw_json -> retorna (data.text, data.words) + """ + mode = os.getenv("STT_OUTPUT_MODE", "api_text").lower() + + if mode != "raw_json": + return str(user_input or ""), None + + try: + payload = json.loads(user_input) if isinstance(user_input, str) else user_input + data = (payload or {}).get("data", {}) + text = str(data.get("text", "") or "") + words = data.get("words", None) + return text, words + except Exception: + return str(user_input or ""), None + + # ==================================================================================== + # Guardrail + + def _lg_guardrail(self, state: PipelineState) -> PipelineState: + raw_input = state.get("user_input", "") + user_input, words = self._normalize_stt_input(raw_input) + + if self.agent_log: + self.agent_log.log_section(f"GUARDRAIL (estágio atual: {self.stage})") + self.agent_log.log("user_input (normalizado)", user_input) + self.agent_log.log("attacks acumulados", str(self.attacks)) + + # Guardrail de confiança (somente raw_json, sem limite) + if words is not None: + conf = sentence_confidence(words) + print("confiança", conf) + if conf < 0.0 or user_input in [None, ""]: + output = "Desculpe, não consegui entender bem. Pode repetir mais devagar, por favor?" + + if self.agent_log: + self.agent_log.log("guardrail_confiança", f"rejeitado (conf={conf})") + self.agent_log.log("output", output) + + # ANTES: subia direto + # AGORA: buffer + self._buffer_langfuse_trace(user_input, output) + + return { + "is_attack": True, + "output": output, + "auto": False, + "stage": state.get("stage", self.stage), + "user_input": user_input, + } + + # Guardrail original (ataque) + #is_attack = self._check_guardrail(user_input) + #print("is_attack", is_attack) + is_attack = False + output = "Desculpe, não entendi. Poderia repetir?" + + if self.attacks > 4: + output = ( + "Notei que temos um problema de comunicação. " + "Irei finalizar a chamada. A TIM agradece a sua atenção" + ) + + if self.agent_log: + self.agent_log.log("guardrail", f"attacks > 4 — encerrando (attacks={self.attacks})") + self.agent_log.log("output", output) + + # buffer + self._buffer_langfuse_trace(user_input, output) + + return { + "is_attack": False, + "output": output, + "auto": False, + "stage": "DONE", + "user_input": user_input, + } + + if is_attack: + self.attacks += 1 + + if self.agent_log: + self.agent_log.log("guardrail", f"ataque detectado (attacks={self.attacks})") + self.agent_log.log("output", output) + + # buffer + self._buffer_langfuse_trace(user_input, output) + + return { + "is_attack": True, + "output": output, + "auto": False, + "stage": state.get("stage", self.stage), + "user_input": user_input, + } + + if self.agent_log: + self.agent_log.log("guardrail", "passou (sem ataque)") + + return {"is_attack": False, "user_input": user_input} + + def _lg_route_from_guardrail(self, state: PipelineState) -> str: + if state.get("is_attack"): + return "end" + + stage = state.get("stage") or self.stage + if stage == "DONE": + return "end" + + return stage + + def _lg_route_after_stage(self, state: PipelineState) -> str: + if state.get("stage") == "DONE": + return "end" + + if state.get("auto"): + return state["stage"] + + return "end" + + # ==================================================================================== + # Stage nodes + + def _get_effective_input(self, state: PipelineState) -> str: + return "###START###" if state.get("auto") else state.get("user_input", "") + + def _lg_presentation(self, state: PipelineState) -> PipelineState: + self.stage = state.get("stage", self.stage) + self.name = self.stage + + text = self._get_effective_input(state) + + if self.agent_log: + self.agent_log.log_section("PRESENTATION") + self.agent_log.log("input", text) + self.agent_log.log_messages("Histórico ANTES do run", self.agent.messages) + + t0 = time.perf_counter() + result = self.agent.run_presentation(text) + elapsed = time.perf_counter() - t0 + + if self.agent_log: + self.agent_log.log("LLM tempo de execução", f"{elapsed:.3f}s") + self.agent_log.log("LLM output", result["output"]) + self.agent_log.log("tool_calls", str(result.get("tools", []))) + self.agent_log.log_agent_state("presentation", self.agent.state) + self.agent_log.log_messages("Histórico DEPOIS do run", self.agent.messages) + + # buffer do output gerado + self._buffer_langfuse_trace(text, result["output"]) + + if self.agent.state.end_conversation: + self.stage = "DONE" + if self.agent_log: + self.agent_log.log_stage_result("presentation", result["output"], "DONE", False) + return {"stage": "DONE", "output": result["output"], "auto": False} + + if self.agent.state.is_target_customer: + self.update_langfuse_auto(text, result["output"]) + self.stage = "argumentation" + self.agent = ArgumentationAgent( + streaming=self.streaming, + prompt_vars=self.prompt_vars, + ) + if self.agent_log: + self.agent_log.log_stage_result("presentation", result["output"], "argumentation", True) + self.agent_log.log_section("TRANSIÇÃO → ARGUMENTATION") + self.agent_log.log_system_prompt("argumentation", self.agent.system_prompt) + return {"stage": "argumentation", "auto": True} + + if not self.prompt_vars["elegibility"]: + self.stage = "DONE" + if self.agent_log: + self.agent_log.log_stage_result("presentation", result["output"], "DONE (sem elegibilidade)", False) + return {"stage": "DONE", "output": result["output"], "auto": False} + + if self.agent_log: + self.agent_log.log_stage_result("presentation", result["output"], self.stage, False) + + return {"stage": self.stage, "output": result["output"], "auto": False} + + def _lg_argumentation(self, state: PipelineState) -> PipelineState: + self.stage = state.get("stage", self.stage) + self.name = self.stage + + if self.agent_log: + self.agent_log.log_section("ARGUMENTATION") + + if state.get("auto"): + result = {} + result['output'] = self.argumentation_start() + text = "###START###" + if self.agent_log: + self.agent_log.log("input (auto)", text) + self.agent_log.log("argumentation_start output", result["output"]) + else: + text = self._get_effective_input(state) + + if self.agent_log: + self.agent_log.log("input", text) + self.agent_log.log_messages("Histórico ANTES do run", self.agent.messages) + + t0 = time.perf_counter() + result = self.agent.run_argumentation(text) + elapsed = time.perf_counter() - t0 + + if self.agent_log: + self.agent_log.log("LLM tempo de execução", f"{elapsed:.3f}s") + self.agent_log.log("LLM output", result["output"]) + self.agent_log.log("tool_calls", str(result.get("tools", []))) + self.agent_log.log_messages("Histórico DEPOIS do run", self.agent.messages) + + if self.agent_log: + self.agent_log.log_agent_state("argumentation", self.agent.state) + + # buffer do output gerado + self._buffer_langfuse_trace(text, result["output"]) + + if self.agent.state.accepted: + self.update_langfuse_auto(text, result["output"]) + self.stage = "data_confirmation" + + self.agent = DataConfirmationAgent( + expected_cpf_last_3=self.prompt_vars["customer"]["NUM_CPF_CNPJ_CLIENTE"], + expected_birth_date=self.prompt_vars["customer"]["DAT_NASCIMENTO"], + streaming=self.streaming, + prompt_vars={"customer": self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"]}, + ) + + if self.agent_log: + self.agent_log.log_stage_result("argumentation", result["output"], "data_confirmation", True) + self.agent_log.log_section("TRANSIÇÃO → DATA_CONFIRMATION") + self.agent_log.log_system_prompt("data_confirmation", self.agent.system_prompt) + + return {"stage": "data_confirmation", "auto": True} + + if self.agent.state.end_conversation: + self.stage = "DONE" + if self.agent_log: + self.agent_log.log_stage_result("argumentation", result["output"], "DONE", False) + return {"stage": "DONE", "output": result["output"], "auto": False} + + if self.agent_log: + self.agent_log.log_stage_result("argumentation", result["output"], self.stage, False) + + return {"stage": self.stage, "output": result["output"], "auto": False} + + def _lg_data_confirmation(self, state: PipelineState) -> PipelineState: + self.stage = state.get("stage", self.stage) + self.name = self.stage + + text = self._get_effective_input(state) + text = text.replace("/", "").replace("-", "").replace(".", "") + + if self.agent_log: + self.agent_log.log_section("DATA_CONFIRMATION") + self.agent_log.log("input", text) + self.agent_log.log_messages("Histórico ANTES do run", self.agent.messages) + + t0 = time.perf_counter() + result = self.agent.run_confirmation(text) + elapsed = time.perf_counter() - t0 + + if self.agent_log: + self.agent_log.log("LLM tempo de execução", f"{elapsed:.3f}s") + self.agent_log.log("LLM output", result["output"]) + self.agent_log.log("tool_calls", str(result.get("tools", []))) + self.agent_log.log_agent_state("data_confirmation", self.agent.state) + self.agent_log.log_messages("Histórico DEPOIS do run", self.agent.messages) + + # buffer do output gerado + self._buffer_langfuse_trace(text, result["output"]) + + if self.agent.state.is_authenticated: + self.update_langfuse_auto(text, result["output"]) + self.stage = "formalization" + + self.agent = FormalizationAgent( + streaming=self.streaming, + prompt_vars=self.prompt_vars, + ) + + if self.agent_log: + self.agent_log.log_stage_result("data_confirmation", result["output"], "formalization", True) + self.agent_log.log_section("TRANSIÇÃO → FORMALIZATION") + self.agent_log.log_system_prompt("formalization", self.agent.system_prompt) + + return {"stage": "formalization", "auto": True} + + if self.agent.state.end_conversation: + self.stage = "DONE" + if self.agent_log: + self.agent_log.log_stage_result("data_confirmation", result["output"], "DONE", False) + return {"stage": "DONE", "output": result["output"], "auto": False} + + if self.agent_log: + self.agent_log.log_stage_result("data_confirmation", result["output"], self.stage, False) + + return {"stage": self.stage, "output": result["output"], "auto": False} + + def _lg_formalization(self, state: PipelineState) -> PipelineState: + self.stage = state.get("stage", self.stage) + self.name = self.stage + + text = self._get_effective_input(state) + + if self.agent_log: + self.agent_log.log_section("FORMALIZATION") + self.agent_log.log("input", text) + self.agent_log.log_messages("Histórico ANTES do run", self.agent.messages) + + t0 = time.perf_counter() + result = self.agent.run_formalization(text) + elapsed = time.perf_counter() - t0 + + if self.agent_log: + self.agent_log.log("LLM tempo de execução", f"{elapsed:.3f}s") + self.agent_log.log("LLM output", result["output"]) + self.agent_log.log("tool_calls", str(result.get("tools", []))) + self.agent_log.log_agent_state("formalization", self.agent.state) + self.agent_log.log_messages("Histórico DEPOIS do run", self.agent.messages) + + # buffer do output gerado + self._buffer_langfuse_trace(text, result["output"]) + + if self.agent.state.end_conversation: + self.stage = "DONE" + if self.agent_log: + self.agent_log.log_stage_result("formalization", result["output"], "DONE (end_conversation)", False) + return {"stage": "DONE", "output": result["output"], "auto": False} + + if self.agent.state.accepted: + self.stage = "DONE" + self.accepted_purchase = True + if self.agent_log: + self.agent_log.log_stage_result("formalization", result["output"], "DONE (compra aceita)", False) + return {"stage": "DONE", "output": result["output"], "auto": False} + + if self.agent.state.declines >= 2: + self.stage = "DONE" + if self.agent_log: + self.agent_log.log_stage_result("formalization", result["output"], f"DONE (declines={self.agent.state.declines})", False) + return {"stage": "DONE", "output": result["output"], "auto": False} + + if self.agent_log: + self.agent_log.log_stage_result("formalization", result["output"], self.stage, False) + + return {"stage": self.stage, "output": result["output"], "auto": False} + + # ==================================================================================== + # Public API + + @timed("Agent run") + def run(self, user_input: Any): + if self.graph is None: + raise RuntimeError("prepare() deve ser chamado antes de run().") + + if self.stage == "DONE": + return "DONE", "" + + if self.agent_log: + self.agent_log.log_section(f"RUN (estágio: {self.stage})") + self.agent_log.log("user_input", str(user_input)) + + #if self.pending_output: + # self.set_interruption(was_interrupted=False) + + initial_state: PipelineState = { + "stage": self.stage, + "user_input": user_input, + "auto": False, + } + + out: PipelineState = self.graph.invoke(initial_state) + + self.stage = out.get("stage", self.stage) + + if self.agent_log: + self.agent_log.log("run finalizado", f"stage={self.stage} output={out.get('output', '')[:120]}...") + + return self.stage, out.get("output", "") + + def _check_guardrail(self, user_input: str): + return self.check_guardrail.run(user_input)["is_attack"] + + # ==================================================================================== + # end_service (mantém compatibilidade com _end_service) + + @timed("Agent end service") + async def end_service(self): + return await self._end_service() + + async def _end_service(self): + if self.pending_output: + self.set_interruption(was_interrupted=False) + + await asyncio.sleep(10) + historical_conversation = get_historic_conversation(self.trace_id) + + self.classified_conversation = self.conversation_classifier.run(historical_conversation) + self.summary = self.conversation_summary.run(historical_conversation) + + success_purchase = {} + + if self.accepted_purchase: + success_purchase = format_success_purchase(self.prompt_vars["raw_customer"]) + + formatted_call = format_call(self.prompt_vars["raw_customer"], self.classified_conversation, self.summary) + output = { + "success_purchase": success_purchase, + "formatted_call": formatted_call, + } + + # Log de fim de serviço + if self.agent_log: + self.agent_log.log_section("END SERVICE") + self.agent_log.log("accepted_purchase", str(self.accepted_purchase)) + self.agent_log.log("classified_conversation", str(self.classified_conversation)) + self.agent_log.log("summary", str(self.summary)) + self.agent_log.log("formatted_call", str(formatted_call)) + self.agent_log.log("success_purchase", str(success_purchase)) + + with langfuse.start_as_current_observation( + as_type="span", + name=self.name, + trace_context={"trace_id": self.trace_id}, + ) as span: + langfuse.update_current_span( + metadata={ + "stage": "classify_conversation", + "classification": formatted_call["MOTIVO_STATUS"], + }, + output=output, + name=self.name, + input=historical_conversation, + ) + + span.score_trace( + name="success_purchase", + value=1 if self.accepted_purchase else 0, + data_type="BOOLEAN", + ) + + return output + + # ==================================================================================== + # Metadata + + def get_formatted_metadata(self): + if self.stage == "presentation": + return { + "stage": self.stage, + "attacks": self.attacks, + "presentation.name_validation_attempt": self.agent.state.name_validation_attempt, + "presentation.eligibility": self.prompt_vars["elegibility"], + "presentation.is_target_customer": self.agent.state.is_target_customer, + "presentation.end_conversation": self.agent.state.end_conversation, + } + + if self.stage == "argumentation": + return { + "stage": self.stage, + "attacks": self.attacks, + "argumentation.accepted": self.agent.state.accepted, + "argumentation.accepted_count": self.agent.state.accepted_count, + "argumentation.declines": self.agent.state.declines, + "argumentation.end_conversation": self.agent.state.end_conversation, + "argumentation.off_topic": self.agent.state.off_topic, + } + + if self.stage == "data_confirmation": + return { + "stage": self.stage, + "attacks": self.attacks, + "data_confirmation.authenticated": self.agent.state.is_authenticated, + "data_confirmation.validated_target_customer": self.agent.state.validated_target_customer, + "data_confirmation.validated_cpf": self.agent.state.validated_cpf, + "data_confirmation.validated_birth_date": self.agent.state.validated_birth_date, + "data_confirmation.attempt_customer": self.agent.state.attempt_customer, + "data_confirmation.attempt_cpf": self.agent.state.attempt_cpf, + "data_confirmation.attempt_birth_date": self.agent.state.attempt_birth_date, + } + + if self.stage == "formalization": + return { + "stage": self.stage, + "attacks": self.attacks, + "formalization.accepted": self.agent.state.accepted, + "formalization.declines": self.agent.state.declines, + "formalization.questions": self.agent.state.questions, + } + + return { + "stage": self.stage, + "attacks": self.attacks, + } diff --git a/src/agent/pipeline/customer_pipeline_unified.py b/src/agent/pipeline/customer_pipeline_unified.py new file mode 100644 index 0000000..0757280 --- /dev/null +++ b/src/agent/pipeline/customer_pipeline_unified.py @@ -0,0 +1,663 @@ +from typing import TypedDict, Any, Optional, Tuple +from langgraph.graph import StateGraph, END + +from langfuse import get_client +from agent.stage.unified_agent import UnifiedAgent +from agent.stage.unified_state import ConversationPhase +from agent.utils.utils import ( + process_mailing, + process_plan, + get_gender, + get_plan, + normalize_name, + sentence_confidence, + start_message_argumentation, +) + +from agent.classifier.check_guardrail import CheckGuardrail +from agent.classifier.conversation_summary import ConversationSummary +from agent.classifier.conversation_classifier import ( + get_historic_conversation, + ConversationClassifier, + format_success_purchase, + format_call, +) + +from app.common.timed import timed +from langfuse.langchain import CallbackHandler +import asyncio +import copy +import os +import json +import math +from app.utils.logging import setup_minimal_logging + +logger = setup_minimal_logging() +langfuse = get_client() +langfuse_handler = CallbackHandler() + + +class PipelineState(TypedDict, total=False): + stage: str + user_input: dict + auto: bool + output: str + is_attack: bool + + +class CustomerPipelineUnified: + """Pipeline unificado: um único agente com prompt dinâmico para toda a conversa.""" + + def __init__(self, mailing, streaming=False): + self.streaming = streaming + self.stage = "presentation" + self.classified_conversation = None + self.accepted_purchase = False + self.name = None + self.attacks = 0 + + self.prompt_vars = { + "raw_customer": mailing, + "customer": None, + "actual_plan": None, + "target_plan": None, + "elegibility": None, + } + + self.trace_id = langfuse.create_trace_id() + self.intro = None + self.graph = None + self.agent: UnifiedAgent | None = None + + # Buffer do último output (aguarda sinal do backend para subir no Langfuse) + self.pending_input = None + self.pending_output = None + self.pending_stage = None + self.pending_metadata = None + print("RODANDO O PIPELINE UNIFICADO") + # ==================================================================================== + # Buffer + commit no Langfuse + # ==================================================================================== + + def _buffer_langfuse_trace(self, input_value: Any, output_value: str): + """ + Guarda o último par (input, output) gerado pelo agente para só subir no Langfuse + quando o backend sinalizar se foi ouvido por completo ou interrompido. + """ + self.pending_input = input_value + self.pending_output = output_value + self.pending_stage = self.stage + self.pending_metadata = self.get_formatted_metadata() + + def update_langfuse_trace(self, final_output: str, was_interrupted: bool = False): + """ + Faz o commit no Langfuse usando os campos pendentes (input/stage/metadata) + e o output final (completo ou truncado). + """ + if self.pending_stage is None: + return + + metadata = self.pending_metadata or {} + metadata = {**metadata, "interrupted": was_interrupted} + + with langfuse.start_as_current_observation( + as_type="span", + name=self.pending_stage, + trace_context={"trace_id": self.trace_id}, + ) as span: + langfuse.update_current_span( + metadata=metadata, + name=self.pending_stage, + output=final_output, + input=self.pending_input, + ) + + def update_langfuse_auto(self, input, output): + with langfuse.start_as_current_observation( + as_type="span", + name=self.stage, + trace_context={"trace_id": self.trace_id}, + ) as span: + langfuse.update_current_span( + metadata=self.get_formatted_metadata(), + name=self.stage, + output=output, + input=input + ) + + # ==================================================================================== + # Backend sinaliza consumo/interrupção (TTS) + # ==================================================================================== + + def set_interruption( + self, + was_interrupted: bool, + listened_text: Optional[str] = None, + skipped_vacalization: bool = False): + + logger.info( + "SET_INTERRUPTION | was_interrupted=%s | listened_text=%r | skipped_vacalization=%s", + was_interrupted, + listened_text, + skipped_vacalization, + ) + + if not self.pending_output: + return + + # Definimos a marca de interrupção + INTERRUPTION_TAG = "###interrupção###" + + if skipped_vacalization: + # Caso 1: O áudio nem chegou a tocar (interrupção imediata ou erro) + final_output = INTERRUPTION_TAG + was_interrupted_flag = True + self.agent.state.restore_state() + self.agent.set_interruption_message(final_output) + + elif was_interrupted: + # Caso 2: Estava falando e foi cortado + # Se ouviu algo, usa o trecho + tag. Se não ouviu nada, apenas a tag. + final_output = f"{listened_text or ''} {INTERRUPTION_TAG}".strip() + was_interrupted_flag = True + self.agent.set_interruption_message(final_output) + + else: + # Caso 3: Fluxo normal (não foi interrompido) + final_output = self.pending_output + was_interrupted_flag = False + + # Atualiza o Rastreamento (Langfuse) + self.update_langfuse_trace( + final_output, + was_interrupted=was_interrupted_flag, + ) + + # Limpa buffer + self.pending_input = None + self.pending_output = None + self.pending_stage = None + self.pending_metadata = None + + # ==================================================================================== + # Start + Prepare + # ==================================================================================== + + def start(self): + first_name = self.prompt_vars["raw_customer"]["NOME_CLIENTE_COMPLETO"].split()[0].lower() + self.intro = ( + f"Olá, meu nome é Helena, sou consultora de vendas da TIM, para sua segurança a ligação está sendo gravada. " + f"Eu estou falando com {first_name}?" + ) + self.stage = "presentation" + self.name = "presentation" + return self.stage, self.intro + + def _generate_argumentation_start_text(self): + """Gera a mensagem fixa de abertura da argumentação.""" + text = start_message_argumentation({ + "cliente_alvo_primeiro_nome": self.prompt_vars["customer"]["cliente_alvo_primeiro_nome"].lower(), + "meses_restantes_fidelizacao": self.prompt_vars["actual_plan"]["meses_restantes_fidelizacao"], + "quanto_pago_a_mais": self.prompt_vars["target_plan"]["quanto_pago_a_mais"], + "quanto_pago_a_mais_int": abs(math.ceil( + self.prompt_vars["raw_customer"]["VLR_FINAL_PLANO_DEST"] - self.prompt_vars["raw_customer"]["MEDIA_RECARGA_2"] + )), + "gb_plano_atual": self.prompt_vars["actual_plan"]["gb_plano_atual"], + "valor_plano_sem_fidelização": self.prompt_vars["actual_plan"]["valor_plano_sem_fidelização"], + "dados_GB": self.prompt_vars["target_plan"]["dados_GB"], + "beneficios_001": self.prompt_vars["target_plan"]["beneficios_001"], + "valor_plano_final": self.prompt_vars["target_plan"]["valor_plano_final"], + "gb_alvo_diferenca": self.prompt_vars["target_plan"]["gb_alvo_diferenca"], + "preco_por_dia": self.prompt_vars["target_plan"]["preço_por_dia"] + }) + return text + + def _generate_formalization_start_text(self): + """Gera o texto obrigatório e literal de abertura da formalização.""" + customer = self.prompt_vars["customer"] + target = self.prompt_vars["target_plan"] + tratamento = customer.get("TRATAMENTO", "Você") + + text = ( + f"Pra finalizar eu preciso apresentar uma série de informações sobre o novo plano. " + f"Vou começar: {tratamento} está adquirindo uma oferta do plano {target['plano']}, " + f"que por doze meses terá o custo de {target['valor_plano_final']}. " + f"Este é um desconto de {target['valor_desconto']} em relação ao valor total do plano, " + f"que é {target['valor_plano_bruto']}. Nesta oferta, {tratamento} tem {target['dados_GB']} " + f"gigas para navegar na internet. Além disso, você tem os seguintes benefícios no novo plano: " + f"{target['beneficios_todos']} e também recebe uma série de serviços de valor agregado. " + f"Você pode conhecê-los no APP meu Tim. A sua data de vencimento e forma de pagamento " + f"permanecem os mesmos. Na sua próxima fatura será cobrado o valor integral do plano antigo, " + f"mais o proporcional dos dias utilizados do plano adquirido. {tratamento} confirma a migração " + f"para o {target['plano']} no valor de {target['valor_plano_final']} fidelizado por doze meses?? " + f"Se sim, diga eu confirmo." + ) + return text + + def _generate_data_confirmation_start_text(self): + """Gera o texto de abertura da confirmação de dados.""" + nome = self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"] + text = ( + f"Ótimo! então vamos ativar agora para você já aproveitar! " + f"Para prosseguirmos você pode por favor confirmar se seu nome é {nome}? " + f"Se sim, diga eu confirmo" + ) + return text + + # ==================================================================================== + # Prepare — instancia o UnifiedAgent + # ==================================================================================== + + def prepare(self, elegibility: bool, protocol: str): + self.stage = "presentation" + self.name = "presentation" + + self.prompt_vars["customer"] = copy.deepcopy(self.prompt_vars["raw_customer"]) + self.prompt_vars["customer"]["PROTOCOLO"] = protocol + + self.prompt_vars["elegibility"] = elegibility + + self.prompt_vars["actual_plan"] = process_plan({}, self.prompt_vars["customer"], "actual") + self.prompt_vars["target_plan"] = process_plan( + get_plan(self.prompt_vars["customer"]["PLANO_DESTINO"]), + self.prompt_vars["customer"], + "target", + ) + + self.prompt_vars["customer"] = get_gender(self.prompt_vars["customer"]) + self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"] = normalize_name( + self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"] + ) + + self.prompt_vars["customer"] = process_mailing(self.prompt_vars["customer"]) + + # ── Instancia o agente unificado ── + self.agent = UnifiedAgent( + expected_cpf_last_3=self.prompt_vars["customer"]["NUM_CPF_CNPJ_CLIENTE"], + expected_birth_date=self.prompt_vars["customer"].get("DAT_NASCIMENTO", ""), + streaming=self.streaming, + prompt_vars={ + "customer_first_name": self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"] + .split()[0].lower(), + "customer_full_name": self.prompt_vars["customer"]["NOME_CLIENTE_COMPLETO"], + "customer": self.prompt_vars["customer"], + "actual_plan": self.prompt_vars["actual_plan"], + "target_plan": self.prompt_vars["target_plan"], + "tratamento": self.prompt_vars["customer"].get("TRATAMENTO", "O senhor"), + "plano_target": self.prompt_vars["target_plan"]["plano"], + "valor_plano_final": self.prompt_vars["target_plan"]["valor_plano_final"], + "valor_desconto": self.prompt_vars["target_plan"]["valor_desconto"], + "valor_plano_bruto": self.prompt_vars["target_plan"]["valor_plano_bruto"], + "dados_GB": self.prompt_vars["target_plan"]["dados_GB"], + "beneficios_todos": self.prompt_vars["target_plan"]["beneficios_todos"], + "elegibility": self.prompt_vars["elegibility"], + }, + ) + + # Injeta o início da conversa no histórico + self.agent.inject_user_message("###START###") + self.agent.inject_ai_message(self.intro) + + self.check_guardrail = CheckGuardrail() + self.conversation_classifier = ConversationClassifier() + self.conversation_summary = ConversationSummary() + + self.update_langfuse_auto( + {"text": "###START###", "prompt_vars": self.prompt_vars}, + self.intro, + ) + + self.graph = self._build_graph() + + # ==================================================================================== + # Grafo simplificado: guardrail → unified → END + # ==================================================================================== + + def _build_graph(self): + g = StateGraph(PipelineState) + + g.add_node("guardrail", self._lg_guardrail) + g.add_node("unified", self._lg_unified) + + g.set_entry_point("guardrail") + + g.add_conditional_edges( + "guardrail", + self._lg_route_from_guardrail, + { + "unified": "unified", + "end": END, + }, + ) + + g.add_conditional_edges( + "unified", + self._lg_route_after_stage, + { + "unified": "unified", + "end": END, + }, + ) + + return g.compile() + + # ==================================================================================== + # STT normalization + # ==================================================================================== + + def _normalize_stt_input(self, user_input: Any) -> Tuple[str, Optional[list]]: + """ + api_text -> retorna (texto, None) + raw_json -> retorna (data.text, data.words) + """ + mode = os.getenv("STT_OUTPUT_MODE", "api_text").lower() + + if mode != "raw_json": + return str(user_input or ""), None + + try: + payload = json.loads(user_input) if isinstance(user_input, str) else user_input + data = (payload or {}).get("data", {}) + text = str(data.get("text", "") or "") + words = data.get("words", None) + return text, words + except Exception: + return str(user_input or ""), None + + # ==================================================================================== + # Guardrail (idêntico ao original) + # ==================================================================================== + + def _lg_guardrail(self, state: PipelineState) -> PipelineState: + raw_input = state.get("user_input", "") + user_input, words = self._normalize_stt_input(raw_input) + + # Guardrail de confiança (somente raw_json, sem limite) + if words is not None: + conf = sentence_confidence(words) + print("confiança", conf) + if conf < 0.0 or user_input in [None, ""]: + output = "Desculpe, não consegui entender bem. Pode repetir mais devagar, por favor?" + + self._buffer_langfuse_trace(user_input, output) + + return { + "is_attack": True, + "output": output, + "auto": False, + "stage": state.get("stage", self.stage), + "user_input": user_input, + } + + # Guardrail original (ataque) + is_attack = self._check_guardrail(user_input) + output = "Desculpe, não entendi. Poderia repetir?" + + if self.attacks > 4: + output = ( + "Notei que temos um problema de comunicação. " + "Irei finalizar a chamada. A TIM agradece a sua atenção" + ) + + self._buffer_langfuse_trace(user_input, output) + + return { + "is_attack": False, + "output": output, + "auto": False, + "stage": "DONE", + "user_input": user_input, + } + + if is_attack: + self.attacks += 1 + + self._buffer_langfuse_trace(user_input, output) + + return { + "is_attack": True, + "output": output, + "auto": False, + "stage": state.get("stage", self.stage), + "user_input": user_input, + } + + return {"is_attack": False, "user_input": user_input} + + def _lg_route_from_guardrail(self, state: PipelineState) -> str: + if state.get("is_attack"): + return "end" + + stage = state.get("stage") or self.stage + if stage == "DONE": + return "end" + + return "unified" + + def _lg_route_after_stage(self, state: PipelineState) -> str: + if state.get("stage") == "DONE": + return "end" + + if state.get("auto"): + return "unified" + + return "end" + + # ==================================================================================== + # Nó unificado — substitui os 4 nós separados + # ==================================================================================== + + def _lg_unified(self, state: PipelineState) -> PipelineState: + """Nó único que delega ao UnifiedAgent e gerencia transições automáticas.""" + self.stage = state.get("stage", self.stage) + self.name = self.agent.state.current_phase.value + + # ── Se é auto-start de uma nova fase, injeta mensagem fixa ── + if state.get("auto"): + phase = self.agent.state.current_phase + + if phase == ConversationPhase.ARGUMENTATION: + # Mensagem fixa de argumentação + text_arg = self._generate_argumentation_start_text() + self.agent.inject_user_message("###START###") + self.agent.inject_ai_message(text_arg) + + # Buffer para Langfuse (transição automática) + self._buffer_langfuse_trace("###START###", text_arg) + + # Retorna a mensagem fixa como output — cliente ouve isso + return { + "stage": phase.value, + "output": text_arg, + "auto": False, + } + + elif phase == ConversationPhase.DATA_CONFIRMATION: + # Roda com ###START### para gerar a mensagem de confirmação + text_dc = self._generate_data_confirmation_start_text() + self.agent.inject_user_message("###START###") + self.agent.inject_ai_message(text_dc) + + self._buffer_langfuse_trace("###START###", text_dc) + + return { + "stage": phase.value, + "output": text_dc, + "auto": False, + } + + elif phase == ConversationPhase.FORMALIZATION: + # Texto obrigatório e literal da formalização + text_form = self._generate_formalization_start_text() + self.agent.inject_user_message("###START###") + self.agent.inject_ai_message(text_form) + + self._buffer_langfuse_trace("###START###", text_form) + + return { + "stage": phase.value, + "output": text_form, + "auto": False, + } + + # ── Fluxo normal: roda o agente unificado ── + text = state.get("user_input", "") + result = self.agent.run_unified(text) + + output = result["output"] + new_phase = result["phase"] + auto = result["auto"] + + # Buffer para Langfuse + self._buffer_langfuse_trace(text, output) + + # Atualiza stage do pipeline + if new_phase == "done": + self.stage = "DONE" + # Verifica se foi aceite final da formalização + if self.agent.state.form_accepted: + self.accepted_purchase = True + return {"stage": "DONE", "output": output, "auto": False} + + self.stage = new_phase + + # Se houve transição automática, loga no Langfuse antes de trocar + if auto: + self.update_langfuse_auto(text, output) + + return { + "stage": new_phase, + "output": output, + "auto": auto, + } + + # ==================================================================================== + # Public API + # ==================================================================================== + + @timed("Agent run") + def run(self, user_input: Any): + if self.graph is None: + raise RuntimeError("prepare() deve ser chamado antes de run().") + + if self.stage == "DONE": + return "DONE", "" + + initial_state: PipelineState = { + "stage": self.stage, + "user_input": user_input, + "auto": False, + } + + out: PipelineState = self.graph.invoke(initial_state) + + self.stage = out.get("stage", self.stage) + return self.stage, out.get("output", "") + + def _check_guardrail(self, user_input: str): + return self.check_guardrail.run(user_input)["is_attack"] + + # ==================================================================================== + # end_service (mantém compatibilidade com _end_service) + # ==================================================================================== + + @timed("Agent end service") + async def end_service(self): + return await self._end_service() + + async def _end_service(self): + if self.pending_output: + self.set_interruption(was_interrupted=False) + + await asyncio.sleep(10) + historical_conversation = get_historic_conversation(self.trace_id) + + self.classified_conversation = self.conversation_classifier.run(historical_conversation) + self.summary = self.conversation_summary.run(historical_conversation) + + success_purchase = {} + + if self.accepted_purchase: + success_purchase = format_success_purchase(self.prompt_vars["raw_customer"]) + + formatted_call = format_call(self.prompt_vars["raw_customer"], self.classified_conversation, self.summary) + output = { + "success_purchase": success_purchase, + "formatted_call": formatted_call, + } + + with langfuse.start_as_current_observation( + as_type="span", + name=self.name, + trace_context={"trace_id": self.trace_id}, + ) as span: + langfuse.update_current_span( + metadata={ + "stage": "classify_conversation", + "classification": formatted_call["MOTIVO_STATUS"], + }, + output=output, + name=self.name, + input=historical_conversation, + ) + + span.score_trace( + name="success_purchase", + value=1 if self.accepted_purchase else 0, + data_type="BOOLEAN", + ) + + return output + + # ==================================================================================== + # Metadata — lê do UnifiedState centralizado + # ==================================================================================== + + def get_formatted_metadata(self): + if self.agent is None: + return {"stage": self.stage, "attacks": self.attacks} + + state = self.agent.state + phase = state.current_phase.value + + base = { + "stage": phase, + "attacks": self.attacks, + "current_phase": phase, + "end_conversation": state.end_conversation, + } + + if phase == "presentation": + base.update({ + "presentation.name_validation_attempt": state.name_validation_attempt, + "presentation.eligibility": self.prompt_vars.get("elegibility"), + "presentation.is_target_customer": state.is_target_customer, + }) + + elif phase == "argumentation": + base.update({ + "argumentation.accepted": state.arg_accepted, + "argumentation.accepted_count": state.arg_accepted_count, + "argumentation.declines": state.arg_declines, + "argumentation.off_topic": state.arg_off_topic, + }) + + elif phase == "data_confirmation": + base.update({ + "data_confirmation.authenticated": state.is_authenticated, + "data_confirmation.validated_target_customer": state.validated_target_customer, + "data_confirmation.validated_cpf": state.validated_cpf, + "data_confirmation.validated_birth_date": state.validated_birth_date, + "data_confirmation.attempt_customer": state.attempt_customer, + "data_confirmation.attempt_cpf": state.attempt_cpf, + "data_confirmation.attempt_birth_date": state.attempt_birth_date, + }) + + elif phase == "formalization": + base.update({ + "formalization.accepted": state.form_accepted, + "formalization.declines": state.form_declines, + "formalization.questions": state.form_questions, + }) + + return base diff --git a/src/agent/pipeline/pipeline_streaming_text.py b/src/agent/pipeline/pipeline_streaming_text.py new file mode 100644 index 0000000..0ee7cfe --- /dev/null +++ b/src/agent/pipeline/pipeline_streaming_text.py @@ -0,0 +1,31 @@ +from typing import TypedDict, Any, Optional, Tuple +from langfuse import get_client +from agent.stage.teste import TestAgent +from app.common.timed import timed +import asyncio +import copy +import os +import json +from typing import Any, Dict, Iterator, Tuple + +class CustomerPipeline: + def __init__(self): + self.agent = TestAgent(streaming=True) + self.stage = "presentation" + + def start(self): + return "presentation", "" + + def prepare(self, elegibility: bool, protocol: str): + pass + + @timed("Agent run") + def run(self, user_input: Any): + return self.agent.run_stream(user_input) + + @timed("Agent end service") + async def end_service(self): + return await self._end_service() + + async def _end_service(self): + return {} diff --git a/src/agent/stage/argumentation.py b/src/agent/stage/argumentation.py new file mode 100644 index 0000000..46e0dfb --- /dev/null +++ b/src/agent/stage/argumentation.py @@ -0,0 +1,132 @@ +# argumentation/argumentation.py +from functools import lru_cache +from importlib import resources +from typing import Literal +from langchain.tools import tool +from agent.base.base_stage import BaseAgent +import re + +import copy + +class ArgumentationState: + def __init__(self): + self.declines = 0 + self.accepted_count = 0 + self.accepted = False + self.end_conversation = False + self.off_topic = 0 + + self._saved_state = None + self.save_state() + + def save_state(self): + self._saved_state = copy.deepcopy(self.__dict__) + self._saved_state.pop("_saved_state", None) + + def restore_state(self): + if self._saved_state is None: + raise RuntimeError("Nenhum estado foi salvo ainda.") + + self.__dict__.update(copy.deepcopy(self._saved_state)) + + def print_state(self): + print(f"declines: {self.declines}") + print(f"accepted_count: {self.accepted_count}") + print(f"accepted: {self.accepted}") + print(f"end_conversation: {self.end_conversation}") + print(f"off_topic: {self.off_topic}") + +class ArgumentationAgent(BaseAgent): + def __init__(self, streaming=False, prompt_vars=None): + + state = ArgumentationState() + + tools = self._build_tools() + super().__init__( + tools=tools, + streaming=streaming, + prompt_vars=prompt_vars, + agent_name='argumentation' + ) + self.state = state + + def _build_tools(self): + @tool + def detect_intention_purchase(intention: Literal["customer_declined_purchase", + "customer_accepted_purchase", + "off_topic_from_sale", + "other"]): + """Classifica a intenção do cliente em relação à oferta de troca de plano. +- customer_accepted_purchase: cliente demonstrou interesse claro em comprar o plano (ex: 'quero', 'pode ser', 'sim', 'topo') +- customer_declined_purchase: cliente recusou comprar o plano (ex: 'não quero', 'caro', 'não preciso', 'deixa esse plano') +- off_topic_from_sale: cliente falou sobre assuntos que não têm relação com a venda (ex: 'vou para a praia', 'vem almoçar') +- other: dúvidas, perguntas ou respostas ambíguas sobre o plano + """ + return _handle_intention(intention) + + def _handle_intention(intention): + print(intention) + #self.state.history.append(intention) + + if intention == "customer_accepted_purchase": + self.state.accepted_count += 1 + if self.state.accepted_count >= 2: + self.state.accepted = True + return "###end_argumentation###" + + return "Peça a segunda confirmação" + + self.state.accepted_count = 0 + + if intention == "customer_declined_purchase": + self.state.declines += 1 + if self.state.declines > 2: + return f"{self.state.declines}º, finalize" + return f"{self.state.declines}º rejeição" + + if intention == "off_topic_from_sale": + self.state.off_topic += 1 + if self.state.off_topic > 3: + return "Finalize a conversa" + + return "Lide com o input do usuario" + + return [detect_intention_purchase] + + @lru_cache + def _load_prompt(self) -> str: + prompt_text = resources.files("agent.prompts").joinpath("argumentation.txt").read_text(encoding="utf-8") + general_rules = resources.files("agent.prompts").joinpath("general_rules.txt").read_text(encoding="utf-8") + knowledge_text = resources.files("agent.prompts").joinpath("knowledge_base.txt").read_text(encoding="utf-8") + + return prompt_text.replace("{general_rules}", general_rules).replace("{knowledge_base}", knowledge_text) + + def set_interruption_message(self, text: str): + if self.messages: + self.messages[-1].content = f"{text}" + + def update_prompt_variable(self) -> None: + pattern = r"(### Informações extras ###)(.*?)(### Fim das informações extras ###)" + + new_content = f"Recusas do cliente:{self.state.declines}" + + print(new_content) + def replacer(match): + return f"{match.group(1)}\n{new_content}\n{match.group(3)}" + + self.agent.steps[1].messages[0].prompt.template = re.sub( + pattern, + replacer, + self.agent.steps[1].messages[0].prompt.template, + flags=re.DOTALL + ) + + def run_argumentation(self, user_input): + #self.update_prompt_variable() + self.state.save_state() + result = self.run(user_input) + #print(self.messages) + if ("tim agradece sua atenção" in result['output'].lower()): + self.state.end_conversation = True + + return result \ No newline at end of file diff --git a/src/agent/stage/data_confirmation.py b/src/agent/stage/data_confirmation.py new file mode 100644 index 0000000..5df83d7 --- /dev/null +++ b/src/agent/stage/data_confirmation.py @@ -0,0 +1,153 @@ +# data_confirmation/data_confirmation.py +from functools import lru_cache +from importlib import resources +from langchain.tools import tool +from agent.base.base_stage import BaseAgent +import copy +#from Levenshtein import distance + +#def levenshtein_similarity(str1: str, str2: str) -> float: +# s1, s2 = str1.strip().lower(), str2.strip().lower() +# dist = distance(s1, s2) +# max_len = max(len(s1), len(s2)) +# if max_len == 0: +# return 100.0 +# return (1 - dist / max_len) * 100 + +class AuthState: + def __init__(self, expected_cpf, + #expected_mother, + expected_birth): + self.cpf_last_3_digits = None + #self.mother_full_name = None + self.birth_date_YYYY_MM_DD = None + + self.expected_cpf_last_3_digits = expected_cpf + #self.expected_mother_full_name = expected_mother + self.expected_birth_date_YYYY_MM_DD = expected_birth + + self.validated_target_customer = False + self.validated_cpf = False + #self.validated_mother_name = False + self.validated_birth_date = False + + self.attempt_customer = 0 + self.attempt_cpf = 0 + self.attempt_birth_date = 0 + #self.attempt_mother_name = 0 + + self.end_conversation = False + + self._saved_state = None + self.save_state() + @property + def is_authenticated(self): + return ( + self.validated_target_customer and + (self.validated_cpf or self.validated_birth_date) + #and self.validated_mother_name + ) + + @property + def attempts_exceeded(self): + """Retorna True se qualquer tentativa for maior que 2.""" + return any([ + self.attempt_customer >= 3, + self.attempt_cpf >= 3, + self.attempt_birth_date >= 3, + #self.attempt_mother_name >= 3, + ]) + + def save_state(self): + self._saved_state = copy.deepcopy(self.__dict__) + self._saved_state.pop("_saved_state", None) + + def restore_state(self): + if self._saved_state is None: + raise RuntimeError("Nenhum estado foi salvo ainda.") + + self.__dict__.update(copy.deepcopy(self._saved_state)) + +class DataConfirmationAgent(BaseAgent): + def __init__( + self, + expected_cpf_last_3, + #expected_mother, + expected_birth_date, + streaming=False, + prompt_vars = None + ): + + self.state = AuthState(expected_cpf_last_3, + #expected_mother, + expected_birth_date) + tools = self._build_tools() + super().__init__( + tools=tools, + streaming=streaming, + prompt_vars=prompt_vars, + agent_name='data_confirmation' + ) + + @lru_cache + def _load_prompt(self) -> str: + prompt_text = resources.files("agent.prompts").joinpath("data_confirmation.txt").read_text(encoding="utf-8") + general_rules = resources.files("agent.prompts").joinpath("general_rules.txt").read_text(encoding="utf-8") + return prompt_text.replace("{general_rules}", general_rules) + + # ----------------- TOOLS ----------------- + def _build_tools(self): + @tool + def detect_name_target_customer(is_target_customer: bool): + """Confirma se o interlocutor é o cliente alvo. Chame após o cliente responder à pergunta de confirmação de nome. +is_target_customer=true: cliente confirmou dizendo 'eu confirmo' +is_target_customer=false: cliente negou ou informou outro nome""" + self.state.validated_target_customer = is_target_customer + self.state.attempt_customer += 1 + return is_target_customer + + #@tool + #def check_full_mother_name(name: str): + # """Captura o nome completo da mãe dito pelo cliente. Chame sempre que responder a pergunta sobre o nome da mãe.""" + # self.state.mother_full_name = name + # self.state.attempt_mother_name += 1 + # self.state.validated_mother_name = True if levenshtein_similarity(name, self.state.expected_mother_full_name) > 80 else False + # return "correct" if self.state.validated_mother_name else "incorrect" + + @tool + def check_birth_date(date_YYYY_MM_DD: str): + """Recebe a data de nascimento informada pelo cliente no formato YYYY-MM-DD. Converta datas por extenso para o formato antes de chamar (ex: '1 de janeiro de 2000' → '2000-01-01').""" + print(date_YYYY_MM_DD) + self.state.birth_date_YYYY_MM_DD = date_YYYY_MM_DD + self.state.attempt_birth_date += 1 + self.state.validated_birth_date = date_YYYY_MM_DD == self.state.expected_birth_date_YYYY_MM_DD + return "correct. Diga apenas ###end_data_confirmation###" if self.state.validated_birth_date else "incorrect" + + @tool + def check_cpf_digits(digits: str): + """Recebe os dígitos do CPF informados pelo cliente como string numérica. Converta números por extenso para dígitos antes de chamar (ex: 'um dois meia seis' → '1266'). Envie todos os dígitos ditos, mesmo que mais de 3.""" + print(digits) + digits = digits[-3:] + self.state.cpf_last_3_digits = digits + self.state.attempt_cpf += 1 + self.state.validated_cpf = digits == self.state.expected_cpf_last_3_digits + return "correct. Diga apenas ###end_data_confirmation###" if self.state.validated_cpf else "incorrect" + + @tool + def purchase_cancellation(end: bool): + """Encerra a chamada quando o cliente desistir claramente da compra. Chame com end=true.""" + print(">>>>",end) + self.state.end_conversation = end + return end + + return [detect_name_target_customer, check_birth_date, check_cpf_digits, purchase_cancellation]# , check_full_mother_name + + def set_interruption_message(self, text: str): + if self.messages: + self.messages[-1].content = f"{text}" + + def run_confirmation(self, user_input: str): + self.state.save_state() + result = self.run(user_input.replace('.','').replace('-','').replace('/','')) + + return result \ No newline at end of file diff --git a/src/agent/stage/formalization.py b/src/agent/stage/formalization.py new file mode 100644 index 0000000..595db40 --- /dev/null +++ b/src/agent/stage/formalization.py @@ -0,0 +1,67 @@ +# formalization/formalization.py +from functools import lru_cache +from importlib import resources +from langchain.tools import tool +from agent.base.base_stage import BaseAgent +from typing import Literal + +class ArgumentationState: + def __init__(self): + self.declines = 0 + self.accepted = False + self.questions = 0 + self.end_conversation = False + +class FormalizationAgent(BaseAgent): + def __init__(self, streaming=False, prompt_vars=None): + state = ArgumentationState() + + tools = self._build_tools() + super().__init__(tools=tools, + streaming=streaming, + prompt_vars=prompt_vars, + agent_name='formalization') + self.state = state + + def _build_tools(self): + @tool + def detect_intention_purchase(intention: Literal["customer_declined_purchase", "customer_accepted_purchase", "customer_asked_question", "other"]): + """Classifica a intenção do cliente durante a formalização da troca de plano. +- customer_accepted_purchase: cliente disse 'eu confirmo' aceitando a migração +- customer_declined_purchase: cliente recusou confirmar a migração (ex: 'não quero', 'mudei de ideia') +- customer_asked_question: cliente fez uma pergunta sobre o plano, benefícios ou cobrança +- other: resposta que não se encaixa nas categorias acima""" + #print(">>>>",intention) + if intention == "customer_declined_purchase": + self.state.declines += 1 + if intention == "customer_accepted_purchase": + self.state.accepted = True + if intention == "customer_asked_question": + self.state.questions += 1 + + return intention + @tool + def purchase_cancellation(end: bool): + """Encerra a chamada quando o cliente desistir claramente da compra. Chame com end=true.""" + #print(">>>>",end) + self.state.end_conversation = end + return end + + return [detect_intention_purchase, purchase_cancellation] + + @lru_cache + def _load_prompt(self) -> str: + prompt_text = resources.files("agent.prompts").joinpath("formalization.txt").read_text(encoding="utf-8") + general_rules = resources.files("agent.prompts").joinpath("general_rules.txt").read_text(encoding="utf-8") + knowledge_text = resources.files("agent.prompts").joinpath("knowledge_base.txt").read_text(encoding="utf-8") + + return prompt_text.replace("{general_rules}", general_rules).replace("{knowledge_base}", knowledge_text) + + def set_interruption_message(self, text: str): + if self.messages: + self.messages[-1].content = f"{text}" + + def run_formalization(self, user_input: str): + result = self.run(user_input) + + return result diff --git a/src/agent/stage/presentation.py b/src/agent/stage/presentation.py new file mode 100644 index 0000000..10c7772 --- /dev/null +++ b/src/agent/stage/presentation.py @@ -0,0 +1,91 @@ +# argumentation/argumentation.py +from functools import lru_cache +from importlib import resources +from typing import Literal +from langchain.tools import tool +from agent.base.base_stage import BaseAgent +import copy + +class PresentationState: + def __init__(self): + self.name_validation_attempt = 0 + self.is_target_customer = False + self.end_conversation = False + + self._saved_state = None + self.save_state() + + def save_state(self): + self._saved_state = copy.deepcopy(self.__dict__) + self._saved_state.pop("_saved_state", None) + + def restore_state(self): + if self._saved_state is None: + raise RuntimeError("Nenhum estado foi salvo ainda.") + + self.__dict__.update(copy.deepcopy(self._saved_state)) + + def print_state(self): + print(f"name_validation_attempt: {self.name_validation_attempt}") + print(f"is_target_customer: {self.is_target_customer}") + print(f"end_conversation: {self.end_conversation}") + +class PresentationAgent(BaseAgent): + def __init__(self, streaming=False, prompt_vars=None): + + state = PresentationState() + + tools = self._build_tools() + super().__init__( + tools=tools, + streaming=streaming, + prompt_vars=prompt_vars, + agent_name='presentation' + ) + self.state = state + + def _build_tools(self): + @tool + def classify_customer_response(response: Literal["confirmed_target_customer", + "denied_target_customer", + "other"] + ): + """Classifica a resposta do interlocutor sobre ser ou não o cliente alvo. +confirmed_target_customer: confirmou ser o cliente alvo (ex: 'sou eu', 'sim', 'pode falar', disse o próprio nome) +denied_target_customer: negou ser o cliente alvo ou é terceiro (ex: 'não é ele', 'sou o filho', 'ele saiu') +other: resposta ambígua ou não relacionada à confirmação de identidade""" + + print(response) + self.state.name_validation_attempt += 1 + + if self.state.name_validation_attempt >= 4: + return "Finalize a conversa com uma mensagem que contenha no final 'A TIM agradece sua atenção.'" + + if response == "confirmed_target_customer": + self.state.is_target_customer = True + return "###end_presentation###" + + return "Lide com o input do usuario" + + return [classify_customer_response] + + @lru_cache + def _load_prompt(self) -> str: + prompt_text = resources.files("agent.prompts").joinpath("presentation.txt").read_text(encoding="utf-8") + general_rules = resources.files("agent.prompts").joinpath("general_rules.txt").read_text(encoding="utf-8") + return prompt_text.replace("{general_rules}", general_rules) + + def set_interruption_message(self, text: str): + if self.messages: + self.messages[-1].content = f"{text}" + + def run_presentation(self, user_input): + + self.state.save_state() + + result = self.run(user_input) + + if "tim agradece sua atenção" in result['output'].lower(): + self.state.end_conversation = True + + return result \ No newline at end of file diff --git a/src/agent/stage/prompt_manager.py b/src/agent/stage/prompt_manager.py new file mode 100644 index 0000000..1b6c336 --- /dev/null +++ b/src/agent/stage/prompt_manager.py @@ -0,0 +1,466 @@ +# stage/prompt_manager.py +from importlib import resources +from functools import lru_cache +from agent.stage.unified_state import UnifiedState, ConversationPhase + + +class PromptManager: + """Monta o system prompt dinamicamente conforme a fase atual da conversa.""" + + def __init__(self, state: UnifiedState, prompt_vars: dict): + self.state = state + self.prompt_vars = prompt_vars + + # ══════════════════════════════════════════════════════════ + # Prompt principal + # ══════════════════════════════════════════════════════════ + + def build_prompt(self) -> str: + """Constrói o prompt completo baseado na fase atual do UnifiedState.""" + parts = [ + self._header(), + self._checkpoint(), + self._active_phase_instructions(), + self._customer_info(), + ] + prompt = "\n\n".join(parts) + # Substitui variáveis de prompt + prompt = self._apply_vars(prompt) + + print(f"\n📋 [PromptManager] Fase atual: {self.state.current_phase.value}") + return prompt + + # ══════════════════════════════════════════════════════════ + # Bloco 1 — Header fixo (identidade + regras gerais) + # ══════════════════════════════════════════════════════════ + + def _header(self) -> str: + return """Seu nome é Helena e você é um assistente de vendas da TIM. Você atende as pessoas via chamada telefônica. + +--- + +## REGRAS GERAIS + +- A conversa deve ser natural, curta e objetiva, evitando repetições mecânicas e respostas fixas. +- Sempre substitua corretamente informações entre colchetes por seus respectivos valores. Ex: [cliente_alvo_primeiro_nome] → joão +- Nunca escreva colchetes na resposta final. +- Seu texto será vocalizado. Gere sempre texto corrido, sem listas, sem negrito, sem parágrafos. A pontuação deve ser clara e natural. +- Escreva todos os números por extenso. +- Caso o cliente traga temas fora de troca de plano da TIM, oriente ligar para asterisco um quatro quatro ou usar o aplicativo Meu TIM. +- NUNCA responda diretamente com JSON. Sempre com texto. +- Mensagens que indicam para repetir ou não entendeu não é necessário chamar tool. +- NÃO gere mensagens do tipo "se precisa de algo mais só chamar". +- Só é necessário chamar a tool uma vez por input. PROIBIDO chamar tool mais de uma vez por input. +- Caso o cliente peça para esperar dizendo: "espera", "aguarde", "vou pegar", "só um minuto", "calma ai" ou algo que indique espera, adapte uma resposta conforme o contexto. + +--- + +## CONTORNO DE ROBOTIZAÇÃO + +É estritamente proibido utilizar tom mecânico, artificial ou que indique comportamento de robô. A fala deve soar como uma pessoa real em uma ligação telefônica. + +- É PROIBIDO iniciar mensagens com palavras automáticas ou previsíveis como "entendo", "claro", "perfeito", "certo", "ok". Mas pode usar precedido, por exemplo: "ah entendi", "hmm". +- Essas expressões podem ser usadas apenas no meio da frase, nunca no início. +- É obrigatório, de forma moderada e natural, inserir marcadores de oralidade como: "tá?", "né?", "hmm", "olha", "então", "ah sim" para simular fala humana. +- Essas expressões nunca devem ser usadas em excesso ou repetidas na mesma frase. +- Evite frases excessivamente formais, simétricas ou com estrutura publicitária. +- Sempre que possível, substitua frases duras, diretas ou artificiais por construções mais humanas, mantendo o mesmo significado e objetivo comercial. +- Usar linguagem de benefícios emocionais: "navegar sem preocupação", "postar à vontade com a família", "viajar tranquilo". + +Exemplos de substituição obrigatória: +"Você possui direito a este plano." → "Esse plano fica disponível pra você agora, tá?" +"O valor do plano é [valor_plano_final]." → "O valor fica [valor_plano_final], tranquilo?" +"Esse plano oferece mais benefícios." → "Ele acaba entregando mais benefícios no dia a dia, né?" +"Vamos prosseguir com a troca de plano?" → "A gente segue com a troca então, tá?" +"Esse serviço está incluso." → "Sim, esse serviço já vem incluso, sem custo a mais, tá?" + +Dicionário de marcadores de oralidade: "ah sim", "sabe", "olha", "então" — use diferentes no começo de frase. + +--- + +## INTERRUPÇÃO DE FALA + +Se a última mensagem do agente terminou com ###interrupção###, isso significa que você foi interrompido durante a fala. + +São duas situações: +1. Continuar o que dizia. +2. Lidar com o assunto da interrupção. + +Exemplo 1: +Agente: Judity, para sua segurança a ligação está sendo gravada tá? Tenho uma novidade ótima pra você. Mais internet ###interrupção### +Cliente: Alô? +Agente: Então, Mais internet no total de sessenta gigas para você navegar sem preocupação pagando apenas trinta e três reais a mais. Interessante para você? + +Exemplo 2: +Agente: Judity, para sua segurança a ligação está sendo gravada tá? Tenho uma novidade ótima pra você. Mais internet ###interrupção### +Cliente: oi eu não quero tá +Agente: não quer ouvir mais? + +- NUNCA escreva ###interrupção### no texto.""" + + # ══════════════════════════════════════════════════════════ + # Bloco 2 — Checkpoint dinâmico + # ══════════════════════════════════════════════════════════ + + def _checkpoint(self) -> str: + phase = self.state.current_phase + lines = ["## STATUS DA CONVERSA"] + + # 1. Apresentação + if phase == ConversationPhase.PRESENTATION: + lines.append(f"[→] 1. Apresentação — Validando nome do cliente (tentativa {self.state.name_validation_attempt}/4)") + elif phase.value in ("argumentation", "data_confirmation", "formalization", "done"): + lines.append("[x] 1. Apresentação — Cliente alvo confirmado ✓") + else: + lines.append("[ ] 1. Apresentação") + + # 2. Argumentação + if phase == ConversationPhase.ARGUMENTATION: + lines.append(f"[→] 2. Argumentação — Convencer o cliente (recusas: {self.state.arg_declines}/3, aceites: {self.state.arg_accepted_count}/2)") + elif phase.value in ("data_confirmation", "formalization", "done"): + lines.append("[x] 2. Argumentação — Cliente aceitou a oferta ✓") + else: + lines.append("[ ] 2. Argumentação") + + # 3. Confirmação de dados + if phase == ConversationPhase.DATA_CONFIRMATION: + name_status = "✓" if self.state.validated_target_customer else "✗" + cpf_status = "✓" if self.state.validated_cpf else "✗" + birth_status = "✓" if self.state.validated_birth_date else "✗" + lines.append(f"[→] 3. Confirmação de dados — Nome: {name_status}, CPF: {cpf_status}, Nascimento: {birth_status}") + elif phase.value in ("formalization", "done"): + lines.append("[x] 3. Confirmação de dados — Dados validados ✓") + else: + lines.append("[ ] 3. Confirmação de dados") + + # 4. Formalização + if phase == ConversationPhase.FORMALIZATION: + lines.append(f"[→] 4. Formalização — Aguardando confirmação final (recusas: {self.state.form_declines}/2)") + elif phase == ConversationPhase.DONE: + if self.state.form_accepted: + lines.append("[x] 4. Formalização — Compra finalizada ✓") + else: + lines.append("[x] 4. Formalização — Conversa encerrada") + else: + lines.append("[ ] 4. Formalização") + + return "\n".join(lines) + + # ══════════════════════════════════════════════════════════ + # Bloco 3 — Instruções ativas por fase + # ══════════════════════════════════════════════════════════ + + def _active_phase_instructions(self) -> str: + phase = self.state.current_phase + + if phase == ConversationPhase.PRESENTATION: + return self._instructions_presentation() + elif phase == ConversationPhase.ARGUMENTATION: + return self._instructions_argumentation() + elif phase == ConversationPhase.DATA_CONFIRMATION: + return self._instructions_data_confirmation() + elif phase == ConversationPhase.FORMALIZATION: + return self._instructions_formalization() + else: + return "## CONVERSA ENCERRADA\nNão responda mais. A conversa acabou." + + # ── Presentation ── + + def _instructions_presentation(self) -> str: + return """## ETAPA ATUAL: APRESENTAÇÃO + +Seu objetivo único nesta etapa é confirmar se o primeiro nome do interlocutor corresponde ao primeiro nome do cliente alvo antes de apresentar qualquer oferta. + +- Nunca ofereça a oferta nesta etapa. +- Nunca peça nome completo ou CPF. +- Não diga ou sugira que ligará novamente ou registrará contato. +- Se o cliente der indicios que vai chamar o dono, apenas aguarde. +- Seja dinâmico, mas tente sempre validar. + +### INÍCIO DA CONVERSA +Após receber ###START### diga: +"Olá, meu nome é Helena, sou consultora de vendas da TIM, para sua segurança a ligação está sendo gravada. Eu estou falando com {customer_first_name}?" + +### INTERPRETAÇÃO DAS RESPOSTAS + +#### CONFIRMAÇÃO +Se o interlocutor confirmar de forma direta ou implícita que é o cliente (frases como: "isso", "sou eu", "pode falar", "pois não", "tá", "sim", "é ele"/"é ela", "falando", ou disser o próprio nome): +- Invoque a tool classify_customer_response com confirmed_target_customer. +- Gere uma resposta curta e natural de transição, como: "Que bom! Então..." + +#### NÃO É O CLIENTE ALVO +Se o interlocutor indicar que não é o cliente ou que é familiar, terceiro ou que o cliente não está disponível: +- Invoque a tool classify_customer_response com denied_target_customer. +- Finalize imediatamente a chamada com uma mensagem de despedida adequada. +Nunca tente validar novamente nesses casos. + +#### DÚVIDAS, PERGUNTAS OU SITUAÇÕES INTERMEDIÁRIAS +- Se perguntarem o motivo da ligação: responda de forma breve "É sobre uma oferta da TIM, mas antes preciso confirmar se estou falando com {customer_first_name}." +- Se o cliente demonstrar que não ouviu ou não entendeu: repita a última frase. +- Se o cliente pedir para aguardar: responda "Claro, estou aguardando." +- Se o cliente falar de assunto fora do contexto 3x sem avanço, finalize a chamada. +- Para qualquer outra resposta, tente conduzir naturalmente para a confirmação do nome. + +### FINALIZAÇÃO +Finalize a chamada quando: +- O interlocutor não é o cliente alvo. +- Pedido para ligar em outro momento. +- Cliente não está disponível. +- Mensagens automáticas ou caixa postal. +- 4 tentativas sem confirmação. + +A mensagem final deve conter obrigatoriamente: "TIM agradece sua atenção." + +### EXEMPLOS + +Exemplo 1: +Agent: Olá, meu nome é Helena, sou consultora de vendas da Tim. Eu estou falando com jose? +User: Não, ele foi para o hospital +Agent: ok, irei fechar a conversa por não conseguir contactar o titular. A TIM agradece sua atenção. + +Exemplo 2: +Agent: Olá, meu nome é Helena, sou consultora de vendas da Tim. Eu estou falando com brenda? +User: O que você quer? +Agent: hmm, eu gostaria de apresentar uma oferta mas antes preciso confirmar se falo com a brenda. +User: pois não +(confirmed_target_customer) +Agent: Que bom! Então tenho uma novidade ótima pra você...""" + + # ── Argumentation ── + + def _instructions_argumentation(self) -> str: + return """## ETAPA ATUAL: ARGUMENTAÇÃO + +Seu objetivo é convencer o cliente a trocar de plano. A primeira mensagem desta etapa já foi enviada pelo sistema (a apresentação da oferta). Você deve continuar a conversa a partir da resposta do cliente. + +### REGRAS DA ARGUMENTAÇÃO +- Sempre tente conduzir o cliente à aceitação da oferta. +- Sempre que citar plano ou "servicos_valor_agregado", mantenha aspas exatamente como fornecidas. +- Não invente benefícios, características ou condições que não estejam explicitamente informadas. +- Quando o cliente perguntar sobre um serviço, ele quer saber se está incluso no preço da oferta. +- A oferta apresentada é sempre mais cara que o plano atual, porém com mais benefícios. +- Se precisar reapresentar a oferta, varie a fraseologia e o foco dos argumentos para evitar repetição. +- Você não pode falar com o cliente em outro momento. Se ele pedir para ligar depois, informe que a oferta só pode ser tratada agora. +- Sua função se limita exclusivamente à troca de plano. +- Você pode apenas oferecer o plano alvo apresentado, nunca invente outro plano. Não pode oferecer descontos, modificar oferta ou citar opções mais baratas. +- Mantenha o tom respeitoso e confiante, sem pressionar com insistência longa. +- MÁXIMO de 2 frases por turno. +- SEMPRE termine seu turno passando a vez pro cliente com uma pergunta curta. +- Adicionar social proof sutil: "Muitos clientes como você já estão aproveitando". + +### ROTEIRO DE ARGUMENTAÇÃO (Turn-taking) +O cliente vai responder ao gancho. Seu objetivo é apresentar os detalhes EM PARTES (micro-interações), sempre terminando com uma pergunta de validação curta (ex: "Faz sentido para você?", "Vamos fechar?", "tem interesse?", "Vamos trocar?" NÃO use o mesmo em sequência). + +Ordem de benefícios (obrigatória): +1° Apps e redes sociais que não consomem gigas +2° Serviços ilimitados (ligações, SMS, WhatsApp) +3° Outros benefícios (Deezer, TIM no Avião, roaming, etc.) + +- 1ª rejeição: Perguntar o motivo do cliente não ter interesse. +- 2ª rejeição: Contra-argumentar o ponto que ele disse. Se ainda negar, finalize. +- 3ª rejeição: Finalizar com "Tudo bem... A TIM agradece sua atenção." + +### CONFIRMAÇÃO DE ACEITE +Se o cliente demonstrar interesse ("quero", "pode ser", "topo", "sim"): +- Invoque detect_intention_purchase com customer_accepted_purchase. +- Pergunte novamente: "vamos fechar?" ou similar. +- Se confirmar de novo: + - Invoque detect_intention_purchase com customer_accepted_purchase. + - Gere uma mensagem de transição natural. +- São necessárias exatamente DUAS confirmações. + +### NEGATIVA +Se o cliente recusar ("não", "não quero", "caro", "tá bom do jeito que tá"): +- Invoque detect_intention_purchase com customer_declined_purchase. + +### OFF-TOPIC +- Se o cliente falar coisas aleatórias: off_topic_from_sale. +- 3 mensagens off-topic insistindo: finalizar. + +### FINALIZAÇÃO +- 3 recusas: finalizar com "Tudo bem... A TIM agradece sua atenção." +- 3 off-topics: finalizar com "Tudo bem... A TIM agradece sua atenção." +- Mensagens de finalização devem vir sozinhas. + +### EXEMPLO +Agente: "hmm Legal. No TIM Black você passa a ter [dados_GB] gigas. É bastante internet pra navegar à vontade. Vamos fechar?" +Cliente: "não sei" (other) +Agente: "sabe, e tem mais: Redes sociais não descontam dessa franquia, e o valor fica [valor_plano_final]. O que acha?" +Cliente: "Não quero" (customer_declined_purchase) +Agente: "E tem o Deezer incluso que muitos clientes já aproveitam, o que acha?" +Cliente: "Gostei" (customer_accepted_purchase) +Agente: "Então, Vamos fechar?" +Cliente: "Sim" (customer_accepted_purchase) +Agente: "Ótimo! Vamos prosseguir então..." + +### INFORMAÇÕES DO PLANO +Plano atual: {actual_plan} +Plano alvo: {target_plan} + +### BASE DE CONHECIMENTO +{knowledge_base}""" + + # ── Data Confirmation ── + + def _instructions_data_confirmation(self) -> str: + return """## ETAPA ATUAL: CONFIRMAÇÃO DE DADOS + +Seu objetivo é verificar a identidade do cliente confirmando nome completo e CPF ou data de nascimento. +- Chame apenas uma tool por turno, nunca duas ao mesmo tempo. +- Você NUNCA deve inventar um input que não foi dito pelo cliente. +- Não chame a tool se não sabe o que enviar ou se não entendeu e pediu para repetir. +- "meia" ou "seis" = 6 + +### 1. VALIDAÇÃO DE NOME + +Ao receber ###START### SEMPRE diga: +"Ótimo! então vamos ativar agora para você já aproveitar! Para prosseguirmos você pode por favor confirmar se seu nome é {customer_full_name}? Se sim, diga eu confirmo" + +#### 1.1 Confirmação +Se o cliente dizer "eu confirmo": +- Invoque detect_name_target_customer com is_target_customer=true + +#### 1.2 Negação +Se o cliente negar: +- Invoque detect_name_target_customer com is_target_customer=false e pergunte novamente + +#### 1.3 Confirmação diferente +Se não for dito "eu confirmo", pergunte se ele confirma. + +#### 1.4 Esperar +Se pedir para esperar, gere resposta educada dizendo que irá aguardar. Quando retornar, volte à confirmação. + +### 2. VALIDAÇÃO DO CPF (APÓS confirmar nome) + +Diga: "Poderia informar os últimos três dígitos do seu CPF?" + +- Se fornecer, chame check_cpf_digits. +- NUNCA INVENTE. +- Se negar, siga para data de nascimento. +- Por mais que seja pedido os últimos três dígitos, certifique de passar todos dígitos que foram ditos. Mas NUNCA peça ele completo. + +Formatação do CPF: +- 123.123.123.45 → dígitos = 12312312345 +- um dois meia oito → dígitos = 1268 +- meia sete oito dois vinte e quatro → dígitos = 678224 +- zero vinte e quatro sete sete sete oito quarenta e sete vinte e três → dígitos = 02477784723 + +#### 2.1 CPF correto → avançar para conclusão +#### 2.2 CPF incorreto → pedir data de nascimento +#### 2.3 Esperar → aguardar e retomar + +### 2B. DATA DE NASCIMENTO (se CPF negado ou errado) + +Diga: "Poderia confirmar sua data de nascimento?" + +- Chame check_birth_date com formato YYYY-MM-DD. +- Em português, "70" = 1970. + +#### Data correta → avançar +#### Data incorreta → pedir novamente + +### 3. CONCLUSÃO +Após tudo validado (nome + CPF ou nascimento), retorne uma mensagem vazia "". + +### 4. CANCELAMENTO +Se o cliente desistir de comprar: +- Invoque purchase_cancellation com end=true. +- Diga: "Tudo bem, não vamos prosseguir com a venda. A TIM agradece sua atenção!" + +### REGRAS CRÍTICAS +- Todos itens precisam ser validados: 1. Nome 2. CPF ou data de nascimento +- NUNCA finalize sem completar a verificação. +- NUNCA peça novamente algo já verificado. +- APENAS chame tool referente ao que foi dito, não ao passado. + +### EXEMPLO +Agent: "Ótimo! Para prosseguirmos, confirme se seu nome é fulano da silva? Se sim, diga eu confirmo" +Humano: "eu confirmo" +Agent: "Joia, poderia informar os últimos três dígitos do seu cpf?" +Humano: "123" +Agent: "" + +### INFORMAÇÕES DO CLIENTE ALVO +Informações do cliente: {customer}""" + + # ── Formalization ── + + def _instructions_formalization(self) -> str: + return """## ETAPA ATUAL: FORMALIZAÇÃO + +Seu objetivo é apresentar as informações do plano e obter a confirmação final do cliente. + +### 1.0 INÍCIO OBRIGATÓRIO + +Ao receber ###START### diga EXATAMENTE (substituindo apenas as variáveis entre colchetes): + +"Pra finalizar eu preciso apresentar uma série de informações sobre o novo plano. Vou começar: {tratamento} está adquirindo uma oferta do plano {plano_target}, que por doze meses terá o custo de {valor_plano_final}. Este é um desconto de {valor_desconto} em relação ao valor total do plano, que é {valor_plano_bruto}. Nesta oferta, {tratamento} tem {dados_GB} gigas para navegar na internet. Além disso, você tem os seguintes benefícios no novo plano: {beneficios_todos} e também recebe uma série de serviços de valor agregado. Você pode conhecê-los no APP meu Tim. A sua data de vencimento e forma de pagamento permanecem os mesmos. Na sua próxima fatura será cobrado o valor integral do plano antigo, mais o proporcional dos dias utilizados do plano adquirido. {tratamento} confirma a migração para o {plano_target} no valor de {valor_plano_final} fidelizado por doze meses?? Se sim, diga eu confirmo." + +### 1.1 Repetir +Só repita se o cliente pedir ou for necessário. + +### 2.0 CONFIRMAÇÃO DE VENDA + +#### 2.1 Se o cliente recusar: +- Diga: "Tem alguma dúvida que eu possa esclarecer??" +- Invoque detect_intention_purchase com customer_declined_purchase. + +#### 2.2 Se o cliente fizer uma pergunta: +- Responda brevemente, sem citar todos benefícios de uma vez. +- Se possível, termine com: "Se confirma em mudar de plano, diga: Eu confirmo". +- Invoque detect_intention_purchase com customer_asked_question. + +##### 2.2.1 Recusa após pergunta: +- Diga: "Tudo bem, {customer_first_name}. A Tim agradece sua atenção. Obrigado." +- Invoque detect_intention_purchase com customer_declined_purchase. + +#### 2.3 Se a pessoa aceitar ("eu confirmo"): +- Diga: "Obrigado, {customer_first_name}! Seu pedido de troca foi iniciado e em até vinte e quatro horas sua migração para o novo plano estará concluída. Entre três a cinco dias {tratamento} receberá uma pesquisa de satisfação. Para Tim cliente satisfeito é nota 9 ou 10. A Tim agradece a sua atenção." +- Invoque detect_intention_purchase com customer_accepted_purchase. + +#### 2.4 Se pedir para repetir: +- Diga de forma resumida. + +### 3. CANCELAMENTO +Se desistir de comprar: +- Invoque purchase_cancellation com end=true. +- Diga: "Tudo bem, não vamos prosseguir com a venda. A TIM agradece sua atenção!" + +### REGRAS +- 2 recusas: finalizar com "Tudo bem. A Tim agradece sua atenção." +- Se perguntas fora de contexto: responda e tente terminar com mensagem positiva sugerindo trocar. +- Você não pode fornecer descontos, a oferta é fixa. + +### INFORMAÇÕES DO PLANO +Informações do cliente: {customer} +Plano atual: {actual_plan} +Plano alvo: {target_plan} + +### BASE DE CONHECIMENTO +{knowledge_base}""" + + # ══════════════════════════════════════════════════════════ + # Bloco 4 — Informações do cliente + # ══════════════════════════════════════════════════════════ + + def _customer_info(self) -> str: + return "## FIM DAS INSTRUÇÕES" + + # ══════════════════════════════════════════════════════════ + # Substituição de variáveis + # ══════════════════════════════════════════════════════════ + + def _apply_vars(self, prompt: str) -> str: + """Substitui placeholders {key} pelos valores de prompt_vars.""" + for key, value in self.prompt_vars.items(): + prompt = prompt.replace("{" + key + "}", str(value)) + return prompt + + # ══════════════════════════════════════════════════════════ + # Carregar knowledge base + # ══════════════════════════════════════════════════════════ + + @staticmethod + @lru_cache + def load_knowledge_base() -> str: + return resources.files("agent.prompts").joinpath("knowledge_base.txt").read_text(encoding="utf-8") diff --git a/src/agent/stage/teste.py b/src/agent/stage/teste.py new file mode 100644 index 0000000..ff1521d --- /dev/null +++ b/src/agent/stage/teste.py @@ -0,0 +1,48 @@ +# teste/teste.py +from functools import lru_cache +from importlib import resources +from typing import Literal +from langchain.tools import tool +from agent.base.base_stage import BaseAgent + +class TestState: + def __init__(self): + self.count = 0 + +class TestAgent(BaseAgent): + def __init__(self, streaming=False, prompt_vars=None): + + state = TestState() + + tools = self._build_tools() + super().__init__( + tools=tools, + streaming=streaming, + prompt_vars=prompt_vars, + agent_name='test' + ) + self.state = state + + def _build_tools(self): + @tool + def classify_customer_response(response: Literal["positive","negative", "other"] + ): + """Classifica a resposta do cliente""" + + print(response) + self.state.count += 1 + + return "Lide com o input do usuario" + + return [classify_customer_response] + + @lru_cache + def _load_prompt(self) -> str: + prompt_text = resources.files("agent.prompts").joinpath("test.txt").read_text(encoding="utf-8") + return prompt_text + + def run(self, user_input): + return super().run(user_input) + + def run_stream(self, user_input): + return super().stream_run(user_input) \ No newline at end of file diff --git a/src/agent/stage/unified_agent.py b/src/agent/stage/unified_agent.py new file mode 100644 index 0000000..eadb8ae --- /dev/null +++ b/src/agent/stage/unified_agent.py @@ -0,0 +1,313 @@ +# stage/unified_agent.py +from functools import lru_cache +from importlib import resources +from typing import Literal +from langchain.tools import tool +from agent.base.base_stage import BaseAgent +from agent.stage.unified_state import UnifiedState, ConversationPhase +from agent.stage.prompt_manager import PromptManager + + +class UnifiedAgent(BaseAgent): + """Agente unificado que gerencia toda a conversa com prompt dinâmico.""" + + def __init__( + self, + expected_cpf_last_3: str = "", + expected_birth_date: str = "", + streaming: bool = False, + prompt_vars: dict = None, + ): + self.state = UnifiedState( + expected_cpf_last_3=expected_cpf_last_3, + expected_birth_date=expected_birth_date, + ) + + self._prompt_vars = prompt_vars or {} + + # Carrega knowledge base e adiciona às variáveis + knowledge = PromptManager.load_knowledge_base() + self._prompt_vars["knowledge_base"] = knowledge + + self.prompt_manager = PromptManager(self.state, self._prompt_vars) + + tools = self._build_tools() + + super().__init__( + tools=tools, + streaming=streaming, + prompt_vars={}, # não usar substituição do BaseAgent — o PromptManager cuida + agent_name="unified", + dynamic_prompt=True, # permite atualizar system_prompt entre chamadas + ) + + # ══════════════════════════════════════════════════════════ + # Overrides do BaseAgent + # ══════════════════════════════════════════════════════════ + + def _load_prompt(self) -> str: + """Retorna o prompt dinâmico gerado pelo PromptManager. + + NOTA: este método é chamado uma vez no __init__ do BaseAgent, + mas o prompt real é atualizado dinamicamente via _get_dynamic_prompt(). + """ + return self.prompt_manager.build_prompt() + + def _get_dynamic_prompt(self) -> str: + """Gera o prompt atualizado a cada invocação do agente.""" + return self.prompt_manager.build_prompt() + + # ══════════════════════════════════════════════════════════ + # Tools — todas as 6 tools dos 4 estágios + # ══════════════════════════════════════════════════════════ + + def _build_tools(self): + + # ── Tool 1: Presentation ── + @tool + def classify_customer_response( + response: Literal["confirmed_target_customer", "denied_target_customer", "other"] + ): + """Classifica a resposta do cliente na etapa de apresentação: +confirmed_target_customer: interlocutor confirmou ser o cliente alvo +denied_target_customer: interlocutor negou ser o cliente alvo +other: outro""" + print(f"[Tool] classify_customer_response: {response}") + self.state.name_validation_attempt += 1 + + if self.state.name_validation_attempt >= 4: + return "Finalize a conversa com uma mensagem que contenha no final 'A TIM agradece sua atenção.'" + + if response == "confirmed_target_customer": + self.state.is_target_customer = True + return "Cliente confirmado. Aguarde a próxima etapa." + + if response == "denied_target_customer": + return "Finalize a conversa com uma mensagem de despedida que contenha 'A TIM agradece sua atenção.'" + + return "Lide com o input do usuario" + + # ── Tool 2: Argumentation + Formalization ── + @tool + def detect_intention_purchase( + intention: Literal[ + "customer_declined_purchase", + "customer_accepted_purchase", + "off_topic_from_sale", + "customer_asked_question", + "other", + ] + ): + """O agente deve chamar esta ferramenta sempre que o cliente falar algo. +- customer_declined_purchase: cliente negou a oferta +- customer_accepted_purchase: cliente aceitou comprar +- off_topic_from_sale: conversas aleatórias sem sentido +- customer_asked_question: cliente fez uma pergunta +- other: outro""" + print(f"[Tool] detect_intention_purchase ({self.state.current_phase.value}): {intention}") + + phase = self.state.current_phase + + # ── ARGUMENTATION ── + if phase == ConversationPhase.ARGUMENTATION: + if intention == "customer_accepted_purchase": + self.state.arg_accepted_count += 1 + if self.state.arg_accepted_count >= 2: + self.state.arg_accepted = True + return "Cliente confirmou a compra. Aguarde a próxima etapa." + return "Peça a segunda confirmação" + + # Reset accepted count se não for aceite + self.state.arg_accepted_count = 0 + + if intention == "customer_declined_purchase": + self.state.arg_declines += 1 + if self.state.arg_declines > 2: + return f"{self.state.arg_declines}º recusa. Finalize a conversa com 'Tudo bem... A TIM agradece sua atenção.'" + return f"{self.state.arg_declines}º rejeição" + + if intention == "off_topic_from_sale": + self.state.arg_off_topic += 1 + if self.state.arg_off_topic > 3: + return "Finalize a conversa com 'Tudo bem... A TIM agradece sua atenção.'" + + return "Lide com o input do usuario" + + # ── FORMALIZATION ── + elif phase == ConversationPhase.FORMALIZATION: + if intention == "customer_accepted_purchase": + self.state.form_accepted = True + return intention + + if intention == "customer_declined_purchase": + self.state.form_declines += 1 + return intention + + if intention == "customer_asked_question": + self.state.form_questions += 1 + return intention + + return intention + + return "Lide com o input do usuario" + + # ── Tool 3: Data Confirmation - Nome ── + @tool + def detect_name_target_customer(is_target_customer: bool): + """Detecta se está falando com o cliente alvo ou não. Chame sempre que responder a pergunta se está falando com o cliente.""" + print(f"[Tool] detect_name_target_customer: {is_target_customer}") + self.state.validated_target_customer = is_target_customer + self.state.attempt_customer += 1 + return is_target_customer + + # ── Tool 4: Data Confirmation - CPF ── + @tool + def check_cpf_digits(digits: str): + """Recebe exatamente os dígitos do CPF completo ditos pelo cliente""" + print(f"[Tool] check_cpf_digits: {digits}") + digits = digits[-3:] + self.state.cpf_last_3_digits = digits + self.state.attempt_cpf += 1 + self.state.validated_cpf = digits == self.state.expected_cpf_last_3_digits + return "correct" if self.state.validated_cpf else "incorrect" + + # ── Tool 5: Data Confirmation - Data Nascimento ── + @tool + def check_birth_date(date_YYYY_MM_DD: str): + """Captura a data de nascimento do cliente (YYYY-MM-DD).""" + print(f"[Tool] check_birth_date: {date_YYYY_MM_DD}") + self.state.birth_date_YYYY_MM_DD = date_YYYY_MM_DD + self.state.attempt_birth_date += 1 + self.state.validated_birth_date = date_YYYY_MM_DD == self.state.expected_birth_date_YYYY_MM_DD + return "correct" if self.state.validated_birth_date else "incorrect" + + # ── Tool 6: Cancelamento ── + @tool + def purchase_cancellation(end: bool): + """Encerra a chamada mediante mensagem clara de desistência da compra""" + print(f"[Tool] purchase_cancellation: {end}") + self.state.end_conversation = end + return end + + return [ + classify_customer_response, + detect_intention_purchase, + detect_name_target_customer, + check_cpf_digits, + check_birth_date, + purchase_cancellation, + ] + + # ══════════════════════════════════════════════════════════ + # Set interruption + # ══════════════════════════════════════════════════════════ + + def set_interruption_message(self, text: str): + if self.messages: + self.messages[-1].content = f"{text}" + + # ══════════════════════════════════════════════════════════ + # Run principal — com transição de fase + # ══════════════════════════════════════════════════════════ + + def run_unified(self, user_input: str): + """Executa um turno da conversa e verifica transições de fase. + + Retorna: + dict com keys: output, phase, auto, transition_text + - output: resposta do agente + - phase: fase atual após o turno + - auto: se houve transição automática (próxima fase precisa de START) + - transition_text: texto fixo a ser empurrado para o output quando há transição + """ + self.state.save_state() + + # Atualiza o system prompt dinamicamente antes de rodar + self.system_prompt = self._get_dynamic_prompt() + + # Limpa formatação de CPF/datas no input + if self.state.current_phase == ConversationPhase.DATA_CONFIRMATION: + user_input = user_input.replace(".", "").replace("-", "").replace("/", "") + + result = self.run(user_input) + output = result["output"] + + # Verifica se precisa finalizar + if "tim agradece sua atenção" in output.lower(): + self.state.end_conversation = True + + # ── Transições de fase ── + phase = self.state.current_phase + + # Presentation → Argumentation + if phase == ConversationPhase.PRESENTATION and self.state.is_target_customer: + self.state.current_phase = ConversationPhase.ARGUMENTATION + print(f"\n🔄 Transição: PRESENTATION → ARGUMENTATION") + return { + "output": output, + "phase": ConversationPhase.ARGUMENTATION.value, + "auto": True, + "transition_text": None, # será gerado pelo pipeline com start_message_argumentation + } + + # Argumentation → DataConfirmation + if phase == ConversationPhase.ARGUMENTATION and self.state.arg_accepted: + self.state.current_phase = ConversationPhase.DATA_CONFIRMATION + print(f"\n🔄 Transição: ARGUMENTATION → DATA_CONFIRMATION") + return { + "output": output, + "phase": ConversationPhase.DATA_CONFIRMATION.value, + "auto": True, + "transition_text": None, # será gerado pelo pipeline com ###START### + } + + # DataConfirmation → Formalization + if phase == ConversationPhase.DATA_CONFIRMATION and self.state.is_authenticated: + self.state.current_phase = ConversationPhase.FORMALIZATION + print(f"\n🔄 Transição: DATA_CONFIRMATION → FORMALIZATION") + return { + "output": output, + "phase": ConversationPhase.FORMALIZATION.value, + "auto": True, + "transition_text": None, # será gerado pelo pipeline com texto obrigatório + } + + # Formalization → DONE (aceite) + if phase == ConversationPhase.FORMALIZATION and self.state.form_accepted: + self.state.current_phase = ConversationPhase.DONE + print(f"\n✅ Transição: FORMALIZATION → DONE (aceite)") + return { + "output": output, + "phase": ConversationPhase.DONE.value, + "auto": False, + "transition_text": None, + } + + # Formalization → DONE (2 recusas) + if phase == ConversationPhase.FORMALIZATION and self.state.form_declines >= 2: + self.state.current_phase = ConversationPhase.DONE + print(f"\n❌ Transição: FORMALIZATION → DONE (recusas)") + return { + "output": output, + "phase": ConversationPhase.DONE.value, + "auto": False, + "transition_text": None, + } + + # End conversation (qualquer fase) + if self.state.end_conversation: + self.state.current_phase = ConversationPhase.DONE + return { + "output": output, + "phase": ConversationPhase.DONE.value, + "auto": False, + "transition_text": None, + } + + # Sem transição — continua na mesma fase + return { + "output": output, + "phase": phase.value, + "auto": False, + "transition_text": None, + } diff --git a/src/agent/stage/unified_state.py b/src/agent/stage/unified_state.py new file mode 100644 index 0000000..bd03a7b --- /dev/null +++ b/src/agent/stage/unified_state.py @@ -0,0 +1,107 @@ +# stage/unified_state.py +import copy +from enum import Enum + + +class ConversationPhase(str, Enum): + PRESENTATION = "presentation" + ARGUMENTATION = "argumentation" + DATA_CONFIRMATION = "data_confirmation" + FORMALIZATION = "formalization" + DONE = "done" + + +class UnifiedState: + """Estado unificado que consolida os 4 estágios da conversa.""" + + def __init__( + self, + expected_cpf_last_3: str = "", + expected_birth_date: str = "", + ): + # ── Global ── + self.current_phase = ConversationPhase.PRESENTATION + self.end_conversation = False + + # ── Presentation ── + self.name_validation_attempt = 0 + self.is_target_customer = False + + # ── Argumentation ── + self.arg_declines = 0 + self.arg_accepted_count = 0 + self.arg_accepted = False + self.arg_off_topic = 0 + + # ── Data Confirmation ── + self.expected_cpf_last_3_digits = expected_cpf_last_3 + self.expected_birth_date_YYYY_MM_DD = expected_birth_date + self.cpf_last_3_digits = None + self.birth_date_YYYY_MM_DD = None + self.validated_target_customer = False + self.validated_cpf = False + self.validated_birth_date = False + self.attempt_customer = 0 + self.attempt_cpf = 0 + self.attempt_birth_date = 0 + + # ── Formalization ── + self.form_declines = 0 + self.form_accepted = False + self.form_questions = 0 + + # ── Backup ── + self._saved_state = None + self.save_state() + + # ── Propriedades ── + + @property + def is_authenticated(self): + return ( + self.validated_target_customer + and (self.validated_cpf or self.validated_birth_date) + ) + + @property + def attempts_exceeded(self): + return any([ + self.attempt_customer >= 3, + self.attempt_cpf >= 3, + self.attempt_birth_date >= 3, + ]) + + # ── Save / Restore ── + + def save_state(self): + state = copy.deepcopy(self.__dict__) + state.pop("_saved_state", None) + self._saved_state = state + + def restore_state(self): + if self._saved_state is None: + raise RuntimeError("Nenhum estado foi salvo ainda.") + self.__dict__.update(copy.deepcopy(self._saved_state)) + + # ── Debug ── + + def print_state(self): + print(f"phase: {self.current_phase.value}") + print(f"end_conversation: {self.end_conversation}") + print(f"--- Presentation ---") + print(f" name_validation_attempt: {self.name_validation_attempt}") + print(f" is_target_customer: {self.is_target_customer}") + print(f"--- Argumentation ---") + print(f" arg_declines: {self.arg_declines}") + print(f" arg_accepted_count: {self.arg_accepted_count}") + print(f" arg_accepted: {self.arg_accepted}") + print(f" arg_off_topic: {self.arg_off_topic}") + print(f"--- Data Confirmation ---") + print(f" validated_target_customer: {self.validated_target_customer}") + print(f" validated_cpf: {self.validated_cpf}") + print(f" validated_birth_date: {self.validated_birth_date}") + print(f" is_authenticated: {self.is_authenticated}") + print(f"--- Formalization ---") + print(f" form_declines: {self.form_declines}") + print(f" form_accepted: {self.form_accepted}") + print(f" form_questions: {self.form_questions}") diff --git a/src/agent/utils/utils.py b/src/agent/utils/utils.py new file mode 100644 index 0000000..2988b91 --- /dev/null +++ b/src/agent/utils/utils.py @@ -0,0 +1,254 @@ +import re +import math +from dataclasses import dataclass +from num2words import num2words +from langchain_openai import ChatOpenAI +from langchain_oci import ChatOCIGenAI +#from langchain_google_genai import ChatGoogleGenerativeAI +import math + +from app.common.timed import timed +from dotenv import load_dotenv +import os +load_dotenv() + +def sentence_confidence(words): + probs = [w["probability"] for w in words if w["probability"] > 0] + if not probs: + return 0.0 + log_mean = sum(math.log(p) for p in probs) / len(probs) + return round(math.exp(log_mean), 4) + +"""llm = ChatOpenAI( + model="openai/gpt-oss-20b", + api_key="fake-key", + base_url="http://10.153.34.154/gpt-oss-20b/v1", + temperature=0.0, + top_p=0.1, + reasoning_effort="low", + )""" + +llm = ChatOCIGenAI( + model_id=os.getenv("OCI_ENDPOINT_ID", ""), + service_endpoint=os.getenv("OCI_ENDPOINT", ""), + compartment_id=os.getenv("OCI_COMPARTMENT_ID", ""), + auth_file_location=os.getenv("OCI_AUTH_FILE_LOCATION", "./config"), + model_kwargs={"temperature": 0.0, + "top_p": 0.1, + "reasoning_effort":"MINIMAL"} + ) + + +target_plans = {"TIM_BLACK_C_LIGHT":{ + "plano": "\"TIM Black Cê Light\"", + "beneficios_001": "tem as principais redes sociais sem consumir seus gigas, e ainda ganha o Deezer pra ouvir suas músicas", + "beneficios_002": "não consome seus gigas quando usar as principais redes sociais, tem acesso ao wifi nos aviões da gol e da latân e tem também o pacote chile de roaming internacional", + "beneficios_todos": """ligações e SMS ilimitados. Instagram, Facebook, X e mensagens de texto no Whatsapp sem consumir seus gigas. Acesso ao wifi nos aviões da gol e da latân. Pacote Chile de roaming internacional.""", + "dependentes": "Não há dependentes", + "forma_pagamento": "Fatura", + "comunicacao_whatsapp": "troca de mensagens é ilimitada, mas áudio e vídeo consomes seus Gigas", + "comunicacao_ligacoes_voz": "Ilimitadas", + "comunicacao_sms": "Ilimitados", + "wifi_no_aviao": "TIM no Avião", + "roaming_internacional": "Pacote Chile", + "servicos_valor_agregado": """ "Aya Audiobooks Premium", "Bancah Premium" mais Jornais", "EXA Segurança Premium", "Aya Ensinah Premium", "EXA Cloud", Busuu, "Fluid Premium" """, + "nao_consome_dos_gigas": "Instagram, Facebook, X", + "aplicativos_inclusos": "Não há", + "aplicativos_a_escolher": "Não há" +}, +"TIM_BLACK_A":{ + "plano": "\"TIM Black AH\"", + "beneficios_001": "tem as principais redes sociais sem consumir seus gigas, e ainda ganha o Deezer pra ouvir suas músicas", + "beneficios_002": "tem o aplicativo de músicas Deezer incluído na mensalidade, não consome seus gigas quando usar as principais redes sociais, tem acesso ao wifi nos aviões da gol e da latân e tem também o pacote américas de roaming internacional", + "beneficios_todos": """aplicativo Deezer. Ligações e SMS ilimitados. Instagram, Facebook e X ilimitados. Além disso, acesso ao wifi nos aviões da gol e latân, juntamente com o pacote Américas de roaming internacional.""", + "categoria_plano": "pós pago", + "dependentes": "Não há dependentes", + "forma_pagamento": "Fatura", + "comunicacao_whatsapp": "ilimitado para texto, audio, video", + "comunicacao_ligacoes_voz": "Ilimitadas", + "comunicacao_sms": "Ilimitados", + "wifi_no_aviao": "TIM no Avião", + "roaming_internacional": "Pacote Américas", + "servicos_valor_agregado": """ "Aya Audiobooks Premium", "Bancah Premium mais Jornais", "EXA Segurança Premium", "Aya Ensinah Premium", "EXA Cloud", Busuu, "Fluid Premium", "Fit Me App", "It Game" """, + "nao_consome_dos_gigas": "Instagram, Facebook, X", + "aplicativos_inclusos": "Deezer", + "aplicativos_a_escolher": "Não há" +}, +"TIM_BLACK_C_HERO":{ + "plano": "\"Tim Black Cê Hero\"", + "beneficios_001": "tem as principais redes sociais sem consumir seus gigas, podendo escolher um aplicativo de streaming entre: Deezer, Disney plus, Max, Prime video e Youtube premium", + "beneficios_002": "não consome seus gigas quando usar as principais redes sociais, tem acesso ao wifi nos aviões da gol e da latân e tem também o pacote américas de roaming internacional.", + "beneficios_todos": """Escolher um aplicativo entre: Deezer, Disney plus, Max, Prime video e Youtube premium. Ligações e SMS ilimitados. Instagram, Facebook e X ilimitados. Além disso, acesso ao wifi nos aviões da gol e latân, juntamente com o pacote Américas de roaming internacional.""", + "categoria_plano": "pós pago", + "dependentes": "Não há dependentes", + "forma_pagamento": "Fatura", + "comunicacao_whatsapp": "ilimitado para texto, audio, video", + "comunicacao_ligacoes_voz": "Ilimitadas", + "comunicacao_sms": "Ilimitados", + "wifi_no_aviao": "TIM no Avião", + "roaming_internacional": "Pacote Américas", + "servicos_valor_agregado": """ "Aya Audiobooks Premium", "Bancah Premium mais Jornais", "EXA Segurança Premium", "Aya Ensinah Premium", "EXA Cloud", Busuu, "Fluid Premium", "Fit Me App", "It Game" """, + "nao_consome_dos_gigas": "Instagram, Facebook, X", + "aplicativos_inclusos": "A escolher", + "aplicativos_a_escolher": "Deezer, Disney plus, Max, Prime video, Youtube premium" +} + +} +@timed("normalize_name") +def normalize_name(name): + x = f"""Normalize o nome abaixo: capitalize e acentue conforme o padrão brasileiro. +Nomes estrangeiros não devem ser acentuados. Retorne APENAS o nome, nada mais. +Exemplos: + +fabio assuncao muller → Fábio Assunção Muller +fatima schmidt → Fátima Schmidt +romulo aragao → Rômulo Aragão +angela maria → Ângela Maria +hortencia → Hortência +EMANUELA → Emanuela +jose dos santos → José dos Santos +lucio fernandez → Lúcio Fernandez +GONCALVES → Gonçalves +JOAO PAULO → João Paulo +--- +Nome: {name} +Resultado:""" + result = llm.invoke(x) + #print("Texto:",result.content) + #print("Metadata:",result.usage_metadata) + #print("*"*100) + return result.content + +def get_gender(mailing): + mailing['SEXO'] = llm.invoke("Retorne APENAS as letras M ou F de acordo com o nome se o sexo é masculino ou feminimo não use ' ou ` retorne apenas uma letra, sendo M ou F:" + mailing['NOME_CLIENTE_COMPLETO']).content + mailing['TRATAMENTO'] = "O senhor" if mailing['SEXO'] == "M" else "A senhora" + return mailing + +def money_to_words(value: float) -> str: + if value is None: + return "" + + reais = int(value) + centavos = int(round((value - reais) * 100)) + + partes = [] + if reais > 0: + partes.append(f"{num2words(reais, lang='pt_BR')} real" if reais == 1 else f"{num2words(reais, lang='pt_BR')} reais") + if centavos > 0: + partes.append(f"{num2words(centavos, lang='pt_BR')} centavo" if centavos == 1 else f"{num2words(centavos, lang='pt_BR')} centavos") + + return " e ".join(partes) if partes else "zero real" + +def calculate_data_bonus(mailing: dict) -> float | None: + gb = eval(mailing['BONUS_DESTINO']) + if mailing['PLANO_DESTINO'] == "TIM_BLACK_A": + gb += 15 + if mailing['PLANO_DESTINO'] == "TIM_BLACK_C_LIGHT": + gb += 20 + + return gb + +def process_plan(plan: dict, mailing: dict, plan_type: str = "actual") -> dict: + + plan.pop("preco_reais", None) + plan.pop("dados_GB", None) + if plan_type == "actual": + valor = mailing.get("MEDIA_RECARGA_2") + fidelization = mailing.get("FIDELIZACAO") or {} + if valor: + plan["valor_pago_ultimos_3_meses"] = money_to_words(valor) + else: + plan["valor_pago_ultimos_3_meses"] = "Não informado" + + plan["plano"] = f"\"{mailing["PLANO_ORIGEM"].strip().title()}\"" + + plan["data_expirar_fidelizacao"] = fidelization.get("final_date") + plan["meses_restantes_fidelizacao"] = fidelization.get("expiration") + plan["valor_plano_sem_fidelização"] = money_to_words(mailing["VALOR_PLANO_CORE"]) + + plan["gb_plano_atual"] = num2words(float(mailing.get("DADOS_CORE").replace(',','.').replace('GB','')), lang='pt_BR') + + elif plan_type == "target": + bruto = mailing.get("VLR_PLANO_DESTINO", 0) or 0 + final = mailing.get("VLR_FINAL_PLANO_DEST", 0) or 0 + desconto = round(bruto - final, 2) + plan['plano'] = plan['plano'].title() + + if mailing.get("MEDIA_RECARGA_2"): + plan['quanto_pago_a_mais'] = money_to_words(abs(math.ceil(final - mailing.get("MEDIA_RECARGA_2")))) + else: + plan['quanto_pago_a_mais'] = "Não informado" + + dados_gb = calculate_data_bonus(mailing) + + gb_atual = float(mailing.get("DADOS_CORE").replace(',','.').replace('GB','')) + gb_alvo_diferenca = dados_gb - gb_atual + + plan.update({ + "valor_plano_bruto": money_to_words(bruto), + "valor_desconto": money_to_words(desconto), + "valor_plano_final": money_to_words(final), + "diferenca_reajuste_fidelizacao": money_to_words(mailing.get("VLR_FINAL_PLANO_DEST", 0) - mailing.get("VALOR_PLANO_CORE", 0)), + "dados_GB": num2words(dados_gb, lang='pt_BR'), + "gb_alvo_diferenca": num2words(gb_alvo_diferenca, lang='pt_BR'), + "preço_por_dia": f"{money_to_words(round(final/30, 2))}" + }) + + else: + raise ValueError("Tipo de plano inválido. Use 'actual' ou 'target'.") + #print(plan) + return plan + +def process_mailing(mailing: dict, campos_desejados=None) -> dict: + """ + Filtra e padroniza os dados do mailing conforme as chaves relevantes. + """ + if campos_desejados is None: + campos_desejados = [ + "NOME_CLIENTE_COMPLETO", "NUM_CPF_CNPJ_CLIENTE", "NUM_TELEFONE", + "DT_NASCIMENTO", "PLANO_ORIGEM", "FLG_TRIPLE_A_OP", "TIPO_MAILING", + "DAT_NASCIMENTO", "NOME_MAE", "SEXO", "TRATAMENTO", "BONUS_DESTINO", + "ELEGIBILIDADE", "PROTOCOLO", "cliente_alvo_primeiro_nome" + ] + + mailing["cliente_alvo_primeiro_nome"] = mailing["NOME_CLIENTE_COMPLETO"].split(" ")[0] + mailing["NUM_CPF_CNPJ_CLIENTE"] = mailing.get("NUM_CPF_CNPJ_CLIENTE", "")[-3:] + return {k: v for k, v in mailing.items() if k in campos_desejados} + +def get_plan(plan: str): + return target_plans[plan] + +def start_message_argumentation(variaveis: dict) -> str: + meses_restantes = variaveis.get("meses_restantes_fidelizacao") + quanto_pago = variaveis.get("quanto_pago_a_mais_int") + + # Definição do texto base conforme regras + if meses_restantes is not None and meses_restantes <= 2: + + if quanto_pago is not None and quanto_pago <= 20: + template = """[cliente_alvo_primeiro_nome]! Você tem hoje um plano controle com [gb_plano_atual] gigas de internet. Em pouco tempo, mais exatamente em [um_mes_ou_dois_meses], acaba o período promocional que você possui e ele será reajustado para [valor_plano_sem_fidelização], que é o valor integral do seu plano sem os descontos. Para que você não tenha que pagar este reajuste sem receber novas vantagens, a TIM aprovou sua migração para o plano TIM Black, onde, ao invés dos [gb_plano_atual] gigas de hoje você terá [dados_GB] Gigas pra usar a internet à vontade e também [beneficios_001]. Este novo plano tem o valor de [valor_plano_final] por mês, sem aumento por doze meses. Uma pequena diferença de [quanto_pago_a_mais] comparando o valor reajustado do seu plano, com muito mais benefícios. Vamos mudar seu plano e aproveitar essa super promoção?""" + else: + template = """[cliente_alvo_primeiro_nome]! Você tem hoje um plano controle com [gb_plano_atual] gigas de internet. Em pouco tempo, mais exatamente em [um_mes_ou_dois_meses], acaba o período promocional que você possui e ele será reajustado para [valor_plano_sem_fidelização], que é o valor integral do seu plano sem os descontos. Para que você não tenha que pagar este reajuste sem receber novas vantagens, a TIM aprovou sua migração para o plano TIM Black, onde, ao invés dos [gb_plano_atual] gigas de hoje você terá [dados_GB] Gigas pra usar a internet à vontade e também [beneficios_001]. Este novo plano tem o valor de [valor_plano_final] por mês, sem aumento por doze meses. Vamos mudar seu plano e aproveitar essa super promoção?""" + + else: + if quanto_pago is not None and quanto_pago <= 20: + template = """[cliente_alvo_primeiro_nome], tenho uma novidade ótima pra você. O plano Black com [gb_alvo_diferenca] gigas a mais para você navegar sem preocupação pagando quase a mesma coisa, [quanto_pago_a_mais] a mais no total de [valor_plano_final] mensais. Vamos aproveitar essa condição?""" + else: + template = """[cliente_alvo_primeiro_nome], tenho uma novidade ótima pra você. O plano Black com [gb_alvo_diferenca] gigas a mais para você navegar sem preocupação juntamente com outros benefícios por cerca de [preco_por_dia] por dia, no total de [valor_plano_final] mensais. Vamos aproveitar essa condição?""" + + # Tratamento especial para "um mês / dois meses" + if meses_restantes in [0,1]: + variaveis["um_mes_ou_dois_meses"] = "um mês" + elif meses_restantes == 2: + variaveis["um_mes_ou_dois_meses"] = "dois meses" + else: + variaveis["um_mes_ou_dois_meses"] = "" + + # Função de substituição automática + def substituir(match): + chave = match.group(1) + return str(variaveis.get(chave, "")) + + mensagem_final = re.sub(r"\[([^\]]+)\]", substituir, template) + + return mensagem_final diff --git a/src/app/agent_entry.py b/src/app/agent_entry.py new file mode 100644 index 0000000..21ebf10 --- /dev/null +++ b/src/app/agent_entry.py @@ -0,0 +1,9 @@ +from __future__ import annotations + +from livekit.agents import cli + +from app.livekit.main import server + + +if __name__ == "__main__": + cli.run_app(server) diff --git a/src/app/bridge_entry.py b/src/app/bridge_entry.py new file mode 100644 index 0000000..b77ad43 --- /dev/null +++ b/src/app/bridge_entry.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +import argparse +import os + +import uvicorn + +from app.ws_gateway.main import app + + +def main() -> None: + parser = argparse.ArgumentParser(description="LiveKit WebSocket Gateway (PCM16 16k <-> LiveKit)") + parser.add_argument("--port", type=int, default=8000) + parser.add_argument("--host", type=str, default="0.0.0.0") + parser.add_argument("--log-level", type=str, default="info") + parser.add_argument("--reload", action="store_true") + args = parser.parse_args() + + uvicorn.run( + app, + host=args.host, + port=args.port, + ws_ping_interval=float(os.getenv("UVICORN_WS_PING_INTERVAL_S", "20")), + ws_ping_timeout=float(os.getenv("UVICORN_WS_PING_TIMEOUT_S", "20")), + ws_max_size=50 * 1024 * 1024, + log_level="info" if args.log_level is None else args.log_level, + reload=args.reload, + ) + + +if __name__ == "__main__": + main() diff --git a/src/app/common/call_config.py b/src/app/common/call_config.py new file mode 100644 index 0000000..a871722 --- /dev/null +++ b/src/app/common/call_config.py @@ -0,0 +1,265 @@ +from __future__ import annotations + +from typing import Any, Dict, Mapping + + +FAKE_AGENT_MIN_RESPONSES = 2 +FAKE_AGENT_MAX_RESPONSES = 10 +FAKE_AGENT_MIN_RESPONSE_CHARS = 40 +FAKE_AGENT_MAX_RESPONSE_CHARS = 180 +FAKE_AGENT_DEFAULT_DELAY_MS = 2500 +FAKE_AGENT_MAX_DELAY_MS = 180000 + + +def _as_str(value: Any) -> str: + if value is None: + return "" + return str(value).strip() + + +def _section(payload: Mapping[str, Any], *keys: str) -> Dict[str, Any]: + for key in keys: + value = payload.get(key) + if isinstance(value, Mapping): + return dict(value) + return {} + + +def _pick_str(payload: Mapping[str, Any], *keys: str) -> str: + for key in keys: + value = payload.get(key) + if value not in (None, ""): + return _as_str(value) + return "" + + +def normalize_call_config(payload: Mapping[str, Any] | None) -> Dict[str, Any]: + if not isinstance(payload, Mapping): + return { + "agent_backend": "", + "stt": {}, + "tts": {}, + "vad": {}, + "vad_logging": {}, + "ws": {}, + "agent_fake": {}, + } + + stt = _section(payload, "stt") + tts = _section(payload, "tts") + vad = _section(payload, "vad") + vad_logging = _section(payload, "vadLogging", "vad_logging") + ws = _section(payload, "ws") + agent_fake = _section(payload, "agentFake", "agent_fake") + return { + "agent_backend": _pick_str(payload, "agentBackend", "agent_backend"), + "stt": { + "provider": _pick_str(stt, "provider"), + "language": _pick_str(stt, "language"), + "api_key": _pick_str(stt, "apiKey", "api_key"), + "initial_prompt": _pick_str(stt, "initialPrompt", "initial_prompt"), + "config_override": _pick_str(stt, "configOverride", "config_override"), + "min_prob_single_word": _pick_str(stt, "minProbSingleWord", "min_prob_single_word"), + "disable_vosk": _pick_str(stt, "disableVosk", "disable_vosk"), + }, + "tts": { + "provider": _pick_str(tts, "provider"), + "voice_id": _pick_str(tts, "voiceId", "voice_id"), + "model_id": _pick_str(tts, "modelId", "model_id"), + "language": _pick_str(tts, "language"), + }, + "vad": { + "min_speech_duration": _pick_str(vad, "minSpeechDuration", "min_speech_duration"), + "activation_threshold": _pick_str(vad, "activationThreshold", "activation_threshold"), + "deactivation_threshold": _pick_str(vad, "deactivationThreshold", "deactivation_threshold"), + "min_silence_duration": _pick_str(vad, "minSilenceDuration", "min_silence_duration"), + "prefix_padding_duration": _pick_str(vad, "prefixPaddingDuration", "prefix_padding_duration"), + "pre_backend_wait_notice_fast_on_vad_pause": _pick_str( + vad, + "preBackendWaitNoticeFastOnVadPause", + "pre_backend_wait_notice_fast_on_vad_pause", + ), + "deferred_interruption_min_audio_ms": ( + _pick_str( + vad, + "deferredInterruptionMinAudioMs", + "deferred_interruption_min_audio_ms", + "DEFERRED_INTERRUPTION_MIN_AUDIO_MS", + ) + or _pick_str( + payload, + "deferredInterruptionMinAudioMs", + "deferred_interruption_min_audio_ms", + "DEFERRED_INTERRUPTION_MIN_AUDIO_MS", + ) + ), + "deferred_interruption_enabled": ( + _pick_str( + vad, + "deferredInterruptionEnabled", + "deferred_interruption_enabled", + "DEFERRED_INTERRUPTION_ENABLED", + ) + or _pick_str( + payload, + "deferredInterruptionEnabled", + "deferred_interruption_enabled", + "DEFERRED_INTERRUPTION_ENABLED", + ) + ), + }, + "vad_logging": { + "log_decisions": _pick_str(vad_logging, "logDecisions", "log_decisions"), + "log_activity": _pick_str(vad_logging, "logActivity", "log_activity"), + "activity_min_probability": _pick_str( + vad_logging, + "activityMinProbability", + "activity_min_probability", + ), + }, + "ws": { + "output_gain": _pick_str(ws, "outputGain", "output_gain"), + "audio_in_backlog_shed_enabled": _pick_str( + ws, + "audioInputBacklogShedEnabled", + "audio_in_backlog_shed_enabled", + ), + "audio_in_backlog_shed_threshold_ms": _pick_str( + ws, + "audioInputBacklogShedThresholdMs", + "audio_in_backlog_shed_threshold_ms", + ), + "audio_in_backlog_shed_keep_ms": _pick_str( + ws, + "audioInputBacklogShedKeepMs", + "audio_in_backlog_shed_keep_ms", + ), + "audio_in_latency_metrics_enabled": _pick_str( + ws, + "audioInputLatencyMetricsEnabled", + "audio_in_latency_metrics_enabled", + ), + "audio_in_latency_alert_ms": _pick_str( + ws, + "audioInputLatencyAlertMs", + "audio_in_latency_alert_ms", + ), + "audio_in_latency_log_interval_s": _pick_str( + ws, + "audioInputLatencyLogIntervalS", + "audio_in_latency_log_interval_s", + ), + "livekit_audio_source_queue_size_ms": _pick_str( + ws, + "livekitAudioSourceQueueSizeMs", + "livekit_audio_source_queue_size_ms", + ), + "livekit_audio_source_clear_on_shed": _pick_str( + ws, + "livekitAudioSourceClearOnShed", + "livekit_audio_source_clear_on_shed", + ), + "audio_in_backlog_energy_shed_enabled": _pick_str( + ws, + "audioInputBacklogEnergyShedEnabled", + "audio_in_backlog_energy_shed_enabled", + ), + "audio_in_backlog_energy_shed_max_excess_ms": _pick_str( + ws, + "audioInputBacklogEnergyShedMaxExcessMs", + "audio_in_backlog_energy_shed_max_excess_ms", + ), + "audio_in_backlog_silence_dbfs": _pick_str( + ws, + "audioInputBacklogSilenceDbfs", + "audio_in_backlog_silence_dbfs", + ), + }, + "agent_fake": { + "delay_ms": _pick_str(agent_fake, "delayMs", "delay_ms"), + "responses": _pick_str(agent_fake, "responses"), + }, + } + + +def resolve_agent_backend_name(call_config: Mapping[str, Any] | None, default_backend: str) -> str: + normalized = normalize_call_config(call_config) + return normalized["agent_backend"] or _as_str(default_backend) or "remote_ws" + + +def resolve_stt_overrides(call_config: Mapping[str, Any] | None) -> Dict[str, str]: + normalized = normalize_call_config(call_config) + return dict(normalized["stt"]) + + +def resolve_tts_overrides(call_config: Mapping[str, Any] | None) -> Dict[str, str]: + normalized = normalize_call_config(call_config) + return dict(normalized["tts"]) + + +def resolve_vad_overrides(call_config: Mapping[str, Any] | None) -> Dict[str, str]: + normalized = normalize_call_config(call_config) + return dict(normalized["vad"]) + + +def resolve_vad_logging_overrides(call_config: Mapping[str, Any] | None) -> Dict[str, str]: + normalized = normalize_call_config(call_config) + return dict(normalized["vad_logging"]) + + +def resolve_ws_overrides(call_config: Mapping[str, Any] | None) -> Dict[str, str]: + normalized = normalize_call_config(call_config) + return dict(normalized["ws"]) + + +def _parse_fake_agent_responses(value: Any) -> list[str]: + raw = _as_str(value) + if not raw: + return [] + + responses = [item.strip() for item in raw.split(";")] + if any(not item for item in responses): + raise ValueError("agentFake.responses nao pode conter itens vazios") + if not FAKE_AGENT_MIN_RESPONSES <= len(responses) <= FAKE_AGENT_MAX_RESPONSES: + raise ValueError( + "agentFake.responses deve conter entre " + f"{FAKE_AGENT_MIN_RESPONSES} e {FAKE_AGENT_MAX_RESPONSES} frases" + ) + for index, response in enumerate(responses, start=1): + if not FAKE_AGENT_MIN_RESPONSE_CHARS <= len(response) <= FAKE_AGENT_MAX_RESPONSE_CHARS: + raise ValueError( + f"agentFake.responses[{index}] deve conter entre " + f"{FAKE_AGENT_MIN_RESPONSE_CHARS} e {FAKE_AGENT_MAX_RESPONSE_CHARS} caracteres" + ) + return responses + + +def resolve_fake_agent_overrides(call_config: Mapping[str, Any] | None) -> Dict[str, Any]: + normalized = normalize_call_config(call_config) + agent_fake = dict(normalized["agent_fake"]) + responses = _parse_fake_agent_responses(agent_fake.get("responses")) + if responses and normalized["agent_backend"].lower() != "remote_ws_fake": + raise ValueError( + "agentFake.responses exige callConfig.agentBackend=remote_ws_fake" + ) + + delay_raw = agent_fake.get("delay_ms") + if delay_raw in (None, ""): + delay_ms = FAKE_AGENT_DEFAULT_DELAY_MS if responses else None + else: + try: + delay_ms = int(str(delay_raw).strip()) + except (TypeError, ValueError) as exc: + raise ValueError("agentFake.delayMs deve ser um inteiro") from exc + if not 0 <= delay_ms <= FAKE_AGENT_MAX_DELAY_MS: + raise ValueError( + f"agentFake.delayMs deve estar entre 0 e {FAKE_AGENT_MAX_DELAY_MS}" + ) + + return {"delay_ms": delay_ms, "responses": responses} + + +def is_fake_agent_stress_test(call_config: Mapping[str, Any] | None) -> bool: + """Identify calls that explicitly opt into the scripted fake agent.""" + overrides = resolve_fake_agent_overrides(call_config) + return bool(overrides.get("responses")) diff --git a/src/app/common/stage.py b/src/app/common/stage.py new file mode 100644 index 0000000..0a3cfe5 --- /dev/null +++ b/src/app/common/stage.py @@ -0,0 +1,18 @@ +from enum import Enum + +class Stages(Enum): + presentation = "Faz perguntas do usuário para saber se está falando com o cliente alvo, pode receber respostas do tipo: sim, não, sou eu, espere um pouco etc." + argumentation = "Conversa com o usuário tentando vender um plano de telefone, possui algumas tentativas de vendas, pode receber respostas do tipo: tipo: sim, não, aceito comprar, não quero, não tenho interesse etc." + data_confirmation = "Confirma os dados do cliente, pode receber nomes, datas e números." + formalization = "Faz a confirmação final da venda, vocalizando um pedindo para o cliente confirmar, pode receber respostas do tipo: sim, não, aceito comprar, não quero, não tenho interesse etc." + + @classmethod + def get_value(cls, key: str = "") -> str: + """ + Retorna o value (descrição) direto a partir do nome da stage. + Ex: Stages.get_value("formalization") + """ + try: + return cls[key].value + except KeyError: + return "" diff --git a/src/app/common/timed.py b/src/app/common/timed.py new file mode 100644 index 0000000..b10e174 --- /dev/null +++ b/src/app/common/timed.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +import functools +import inspect +import logging +import time +from collections.abc import Callable + + +def timed( + name: str | None = None, + *, + log_fn: Callable[..., None] | None = None, + logger: logging.Logger | None = None, + unit: str = "s", # "ms" | "s" | "us" + warn_after_s: float = 5.0, +): + """ + Mede o tempo de execução e registra via: + - log_fn(msg), se fornecido + - senão logger.info(msg), usando por padrão o logger "agent_internal_stt" + + Se dt > warn_after_s, emite WARNING e adiciona "(FUNÇÃO LENTA)". + """ + def deco(fn): + label = name or fn.__qualname__ + base_logger = logger or logging.getLogger("agent_internal_stt") + + def fmt(dt: float) -> str: + if unit == "us": + return f"{dt * 1_000_000:.0f} µs" + if unit == "s": + return f"{dt:.6f} s" + return f"{dt * 1000:.2f} ms" + + def emit_timing(dt: float) -> None: + slow = dt > warn_after_s + msg = f"[timed] {label} levou {fmt(dt)}" + (" (FUNÇÃO LENTA)" if slow else "") + + # Se não passou log_fn, usamos o logger e aí dá pra logar warning de verdade + if log_fn is None: + if slow: + base_logger.warning(msg) + else: + base_logger.info(msg) + return + + # Se passou log_fn, tentamos suportar nível sem quebrar compatibilidade + try: + # se o log_fn aceitar algo como log_fn(msg, level="WARNING") + if slow: + log_fn(msg, level="WARNING") + else: + log_fn(msg, level="INFO") + except TypeError: + # fallback: mantém assinatura antiga log_fn(msg) + if slow: + log_fn(f"[WARN] {msg}") + else: + log_fn(msg) + + if inspect.iscoroutinefunction(fn): + @functools.wraps(fn) + async def aw(*args, **kwargs): + t0 = time.perf_counter() + try: + return await fn(*args, **kwargs) + finally: + emit_timing(time.perf_counter() - t0) + return aw + + @functools.wraps(fn) + def w(*args, **kwargs): + t0 = time.perf_counter() + try: + return fn(*args, **kwargs) + finally: + emit_timing(time.perf_counter() - t0) + return w + + return deco \ No newline at end of file diff --git a/src/app/livekit/adapters/agent_backend.py b/src/app/livekit/adapters/agent_backend.py new file mode 100644 index 0000000..87b8605 --- /dev/null +++ b/src/app/livekit/adapters/agent_backend.py @@ -0,0 +1,43 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Protocol, runtime_checkable + + +@dataclass(frozen=True, slots=True) +class BackendReply: + stage: str + text: str = "" + done: bool = False + export_payload: Any = None + metadata: Any = None + + +@runtime_checkable +class AgentBackend(Protocol): + async def prepare( + self, + elegibility: bool, + protocol: str, + ) -> None: ... + + async def run(self, user_input: Any) -> BackendReply: ... + + async def set_interruption( + self, + interrupted: bool, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: ... + + async def set_processing_interruption( + self, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: ... + + async def inject_idle_nudge(self, nudge_text: str) -> None: ... + + async def end_service_once(self) -> BackendReply: ... diff --git a/src/app/livekit/adapters/audio_gain.py b/src/app/livekit/adapters/audio_gain.py new file mode 100644 index 0000000..4198e74 --- /dev/null +++ b/src/app/livekit/adapters/audio_gain.py @@ -0,0 +1,104 @@ +"""Correção de nível (loudness) na saída do TTS, no worker. + +O TTS costuma sair num nível baixo demais para telefonia. Em vez de multiplicar +o áudio às cegas no fim do pipeline (o que satura/clipa nos picos), aqui aplica-se +um makeup gain com um limiter soft-knee (tanh) já na saída do TTS — antes de +publicar no room. Para níveis normais de fala o ganho é praticamente linear; nos +picos ele comprime suavemente em direção ao teto, sem hard clipping. + + y = ceiling * tanh(gain * x / ceiling) + +Config por ambiente: + TTS_OUTPUT_GAIN makeup linear (1.0 = desligado; default 1.0) + TTS_OUTPUT_CEILING_DBFS teto do limiter em dBFS (default -1.0) +""" + +from __future__ import annotations + +import math +import os +from typing import Any, Optional + +import numpy as np + + +def _safe_float(value: Optional[str], default: float) -> float: + if value is None or not str(value).strip(): + return default + try: + parsed = float(value) + except (TypeError, ValueError): + return default + + if not math.isfinite(parsed): + return default + + return parsed + + +def _dbfs_to_linear(dbfs: float) -> float: + return float(10.0 ** (dbfs / 20.0)) + + +class SoftClipGain: + """Makeup gain + limiter soft-knee (tanh) para PCM16 mono little-endian. + + ``gain`` é o ganho linear aplicado à fala em nível normal; ``ceiling`` é o + teto (0..1 do fundo de escala) que os picos nunca ultrapassam. + """ + + def __init__(self, gain: float, ceiling: float = 0.891) -> None: + try: + normalized_gain = float(gain) + except (TypeError, ValueError): + normalized_gain = 1.0 + self.gain = ( + min(10.0, max(0.0, normalized_gain)) + if math.isfinite(normalized_gain) + else 1.0 + ) + self.ceiling = float(min(1.0, max(0.05, ceiling))) + + @property + def enabled(self) -> bool: + return self.gain != 1.0 + + def process(self, pcm: bytes) -> bytes: + if not self.enabled or not pcm: + return pcm + + count = len(pcm) // 2 + if count == 0: + return pcm + + x = np.frombuffer(pcm, dtype=" SoftClipGain: + gain = min(10.0, max(0.0, _safe_float(os.getenv("TTS_OUTPUT_GAIN"), 1.0))) + ceiling_dbfs = min( + 0.0, max(-60.0, _safe_float(os.getenv("TTS_OUTPUT_CEILING_DBFS"), -1.0)) + ) + return SoftClipGain(gain=gain, ceiling=_dbfs_to_linear(ceiling_dbfs)) + + +class GainEmitter: + """Proxy sobre um ``tts.AudioEmitter`` que aplica ``SoftClipGain`` em cada + ``push()``. Todo o resto (initialize/flush/start_segment/end_segment/...) é + encaminhado sem alteração para o emitter real. + """ + + def __init__(self, inner: Any, gain: SoftClipGain) -> None: + self._inner = inner + self._gain = gain + + def push(self, data: bytes) -> None: + self._inner.push(self._gain.process(data)) + + def __getattr__(self, name: str) -> Any: + return getattr(self._inner, name) diff --git a/src/app/livekit/adapters/azure_rest_tts.py b/src/app/livekit/adapters/azure_rest_tts.py new file mode 100644 index 0000000..f6289e4 --- /dev/null +++ b/src/app/livekit/adapters/azure_rest_tts.py @@ -0,0 +1,242 @@ +from __future__ import annotations + +import asyncio +from typing import List +from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit +from xml.sax.saxutils import escape, quoteattr + +import httpx +from livekit.agents import ( + APIConnectionError, + APIStatusError, + APITimeoutError, + tts, + utils, +) +from livekit.agents.types import APIConnectOptions, DEFAULT_API_CONNECT_OPTIONS + + +AZURE_OUTPUT_FORMATS = { + 8000: "raw-8khz-16bit-mono-pcm", + 16000: "raw-16khz-16bit-mono-pcm", + 22050: "raw-22050hz-16bit-mono-pcm", + 24000: "raw-24khz-16bit-mono-pcm", + 44100: "raw-44100hz-16bit-mono-pcm", + 48000: "raw-48khz-16bit-mono-pcm", +} + + +def _append_query_param(url: str, key: str, value: str) -> str: + if not value: + return url + + parsed = urlsplit(url) + query_pairs = parse_qsl(parsed.query, keep_blank_values=True) + if any(existing_key == key for existing_key, _ in query_pairs): + return url + + query_pairs.append((key, value)) + return urlunsplit(parsed._replace(query=urlencode(query_pairs))) + + +def _is_custom_domain_endpoint(url: str) -> bool: + parsed = urlsplit((url or "").strip()) + host = (parsed.hostname or "").strip().lower() + return host.endswith(".cognitiveservices.azure.com") + + +def _prepend_service_prefix(url: str, service: str) -> str: + parsed = urlsplit((url or "").strip()) + path = (parsed.path or "").lstrip("/") + service = (service or "").strip().strip("/") + if not service or path.startswith(f"{service}/"): + return url.rstrip("/") + return urlunsplit(parsed._replace(path=f"/{service}/{path}")).rstrip("/") + + +class AzureRESTTTS(tts.TTS): + def __init__( + self, + *, + voice: str, + language: str | None = None, + sample_rate: int = 16000, + speech_key: str | None = None, + speech_region: str | None = None, + speech_endpoint: str | None = None, + deployment_id: str | None = None, + speech_auth_token: str | None = None, + user_agent: str = "tia-azure-tts/1.0", + timeout_s: float = 30.0, + ) -> None: + super().__init__( + capabilities=tts.TTSCapabilities(streaming=False, aligned_transcript=False), + sample_rate=sample_rate, + num_channels=1, + ) + if sample_rate not in AZURE_OUTPUT_FORMATS: + raise ValueError( + f"Unsupported sample rate {sample_rate}. Supported: {sorted(AZURE_OUTPUT_FORMATS)}" + ) + if not (voice or "").strip(): + raise ValueError("voice is required") + if not ((speech_key or "").strip() or (speech_auth_token or "").strip()): + raise ValueError("speech_key or speech_auth_token is required") + if not ((speech_region or "").strip() or (speech_endpoint or "").strip()): + raise ValueError("speech_region or speech_endpoint is required") + + self._voice = voice.strip() + self._language = (language or "").strip() or None + self._speech_key = (speech_key or "").strip() or None + self._speech_region = (speech_region or "").strip() or None + self._speech_endpoint = (speech_endpoint or "").strip().rstrip("/") or None + self._deployment_id = (deployment_id or "").strip() or None + self._speech_auth_token = (speech_auth_token or "").strip() or None + self._user_agent = (user_agent or "tia-azure-tts/1.0").strip() + self._timeout_s = float(timeout_s) + + @property + def model(self) -> str: + return self._deployment_id or self._voice + + @property + def provider(self) -> str: + return "azure" + + def synthesize( + self, + text: str, + *, + conn_options: APIConnectOptions = DEFAULT_API_CONNECT_OPTIONS, + ) -> "ChunkedStream": + return ChunkedStream(tts=self, input_text=text, conn_options=conn_options) + + def _base_endpoint(self) -> str: + if self._speech_endpoint: + return self._speech_endpoint + assert self._speech_region + service = "voice" if self._deployment_id else "tts" + return f"https://{self._speech_region}.{service}.speech.microsoft.com/cognitiveservices/v1" + + def _request_endpoints(self) -> List[str]: + base = _append_query_param(self._base_endpoint(), "deploymentId", self._deployment_id or "") + endpoints = [base] + + parsed = urlsplit(base) + path = (parsed.path or "").rstrip("/") + if _is_custom_domain_endpoint(base) and path == "/cognitiveservices/v1": + service = "voice" if self._deployment_id else "tts" + endpoints.append( + _append_query_param( + _prepend_service_prefix(base, service), + "deploymentId", + self._deployment_id or "", + ) + ) + return endpoints + + def _build_ssml(self, text: str) -> bytes: + language = self._language or "pt-BR" + escaped_text = escape((text or "").strip()) + return ( + f"" + f"{escaped_text}" + f"" + ).encode("utf-8") + + def _headers(self) -> dict[str, str]: + headers = { + "Content-Type": "application/ssml+xml", + "X-Microsoft-OutputFormat": AZURE_OUTPUT_FORMATS[self.sample_rate], + "User-Agent": self._user_agent, + } + if self._speech_auth_token: + headers["Authorization"] = f"Bearer {self._speech_auth_token}" + elif self._speech_key: + headers["Ocp-Apim-Subscription-Key"] = self._speech_key + return headers + + def synthesize_pcm(self, text: str) -> bytes: + text = (text or "").strip() + if not text: + return b"" + + headers = self._headers() + body = self._build_ssml(text) + last_status_error: httpx.HTTPStatusError | None = None + endpoints = self._request_endpoints() + + with httpx.Client(timeout=self._timeout_s, follow_redirects=True) as client: + for index, endpoint in enumerate(endpoints): + try: + response = client.post(endpoint, headers=headers, content=body) + response.raise_for_status() + if not response.content: + raise RuntimeError("Azure TTS returned empty audio.") + self._speech_endpoint = endpoint + return response.content + except httpx.HTTPStatusError as exc: + should_try_next = ( + exc.response is not None + and exc.response.status_code == 404 + and index < len(endpoints) - 1 + ) + if should_try_next: + last_status_error = exc + continue + raise + + if last_status_error is not None: + raise last_status_error + raise RuntimeError("Azure TTS returned no response.") + + async def aclose(self) -> None: + return None + + +class ChunkedStream(tts.ChunkedStream): + def __init__( + self, + *, + tts: AzureRESTTTS, + input_text: str, + conn_options: APIConnectOptions, + ) -> None: + super().__init__(tts=tts, input_text=input_text, conn_options=conn_options) + self._tts: AzureRESTTTS = tts + + async def _run(self, output_emitter: tts.AudioEmitter) -> None: + output_emitter.initialize( + request_id=utils.shortuuid(), + sample_rate=self._tts.sample_rate, + num_channels=self._tts.num_channels, + mime_type="audio/pcm", + ) + + try: + audio = await asyncio.to_thread(self._tts.synthesize_pcm, self._input_text) + except httpx.TimeoutException as exc: + raise APITimeoutError() from exc + except httpx.HTTPStatusError as exc: + status_code = exc.response.status_code if exc.response is not None else -1 + request_id = exc.response.headers.get("X-RequestId") if exc.response is not None else None + body = exc.response.text if exc.response is not None else None + message = "Azure TTS request failed." + if body: + message = f"{message} {body}" + raise APIStatusError( + message=message, + status_code=status_code, + request_id=request_id, + body=body, + ) from exc + except httpx.RequestError as exc: + raise APIConnectionError(f"Could not connect to Azure TTS: {exc}") from exc + except RuntimeError as exc: + raise APIConnectionError(str(exc), retryable=False) from exc + + output_emitter.push(audio) + output_emitter.flush() diff --git a/src/app/livekit/adapters/backend_factory.py b/src/app/livekit/adapters/backend_factory.py new file mode 100644 index 0000000..27a3c9d --- /dev/null +++ b/src/app/livekit/adapters/backend_factory.py @@ -0,0 +1,135 @@ +from __future__ import annotations + +import os +from typing import Any, Dict + +from app.livekit.adapters.agent_backend import AgentBackend + + +def _env_first(*names: str) -> str: + for name in names: + value = (os.getenv(name, "") or "").strip() + if value: + return value + return "" + + +def _normalize_agent_name(value: Any) -> str: + raw = str(value or "").strip().lower() + aliases = { + "conta": "conta", + "contas": "conta", + "ofert": "oferta", + "oferta": "oferta", + "ofertas": "oferta", + "cobra": "cobranca", + "cobranca": "cobranca", + "cobrança": "cobranca", + "cobrancas": "cobranca", + "cobranças": "cobranca", + } + return aliases.get(raw, raw) + + +def _current_agent_name(remote_agent_context: Dict[str, Any] | None) -> str: + context = remote_agent_context or {} + return _normalize_agent_name( + context.get("agent") + or context.get("Agent") + or context.get("agente") + or "" + ) + + +def build_agent_backend( + *, + intro: str, + backend_name: str | None = None, + remote_agent_context: Dict[str, Any] | None = None, + timeline: Any | None = None, + streaming: bool = False, +) -> AgentBackend: + requested_backend = (backend_name or os.getenv("AGENT_BACKEND", "remote_ws") or "remote_ws").strip().lower() + backend = requested_backend + use_fake_remote = backend in {"remote_ws_fake", "fake_remote_ws", "ws_fake"} + + if use_fake_remote: + from app.livekit.adapters.fake_remote_ws_adapter import FakeRemoteWSAdapter + + fake_delay_ms = (remote_agent_context or {}).get("_fake_agent_delay_ms") + fake_responses = (remote_agent_context or {}).get("_fake_agent_responses") + return FakeRemoteWSAdapter( + intro=intro, + backend_label="remote_ws_fake", + request_context=remote_agent_context or {}, + timeline=timeline, + default_stage=(os.getenv("REMOTE_AGENT_WS_DEFAULT_STAGE", "PRESENTATION") or "PRESENTATION").strip(), + delay_ms=fake_delay_ms, + responses=fake_responses, + ) + + if backend == "langgraph": + raise RuntimeError("AGENT_BACKEND=langgraph is no longer supported in the websocket runtime") + + if backend in {"remote_ws", "ws", "websocket"}: + from app.livekit.adapters.remote_agent_ws_adapter import RemoteAgentWSAdapter + + url = _env_first("REMOTE_AGENT_WS_URL") + if not url: + raise RuntimeError("REMOTE_AGENT_WS_URL must be set when AGENT_BACKEND=remote_ws") + + return RemoteAgentWSAdapter( + intro=intro, + url=url, + backend_label="remote_ws_fake" if use_fake_remote else requested_backend, + url_by_agent={}, + request_context=remote_agent_context or {}, + timeline=timeline, + streaming=streaming, + default_stage=(os.getenv("REMOTE_AGENT_WS_DEFAULT_STAGE", "PRESENTATION") or "PRESENTATION").strip(), + open_timeout_s=float(os.getenv("REMOTE_AGENT_WS_OPEN_TIMEOUT_S", "10")), + read_timeout_s=float(os.getenv("REMOTE_AGENT_WS_READ_TIMEOUT_S", "45")), + write_timeout_s=float(os.getenv("REMOTE_AGENT_WS_WRITE_TIMEOUT_S", "10")), + close_timeout_s=float(os.getenv("REMOTE_AGENT_WS_CLOSE_TIMEOUT_S", "10")), + max_message_bytes=int(os.getenv("REMOTE_AGENT_WS_MAX_MESSAGE_BYTES", str(1024 * 1024))), + ) + + if backend in {"remote_sse", "sse"}: + from app.livekit.adapters.remote_agent_sse_adapter import RemoteAgentSSEAdapter + + url_by_agent = { + "conta": _env_first("REMOTE_AGENT_SSE_URL_CONTA", "REMOTE_AGENT_SSE_URL_CONTAS"), + "oferta": _env_first("REMOTE_AGENT_SSE_URL_OFERTA", "REMOTE_AGENT_SSE_URL_OFERTAS"), + "cobranca": _env_first( + "REMOTE_AGENT_SSE_URL_COBRANCA", + "REMOTE_AGENT_SSE_URL_COBRANCAS", + "REMOTE_AGENT_SSE_URL_COBRA", + ), + } + current_agent = _current_agent_name(remote_agent_context) + url = url_by_agent.get(current_agent, "") + if not url: + expected_env = { + "conta": "REMOTE_AGENT_SSE_URL_CONTA", + "oferta": "REMOTE_AGENT_SSE_URL_OFERTA", + "cobranca": "REMOTE_AGENT_SSE_URL_COBRANCA", + }.get(current_agent, "REMOTE_AGENT_SSE_URL_") + raise RuntimeError( + f"{expected_env} must be set when AGENT_BACKEND=remote_sse" + ) + + return RemoteAgentSSEAdapter( + intro=intro, + url=url, + backend_label=requested_backend, + url_by_agent=url_by_agent, + request_context=remote_agent_context or {}, + timeline=timeline, + streaming=streaming, + default_stage=(os.getenv("REMOTE_AGENT_SSE_DEFAULT_STAGE", "PRESENTATION") or "PRESENTATION").strip(), + connect_timeout_s=float(os.getenv("REMOTE_AGENT_SSE_CONNECT_TIMEOUT_S", "10")), + read_timeout_s=float(os.getenv("REMOTE_AGENT_SSE_READ_TIMEOUT_S", "45")), + write_timeout_s=float(os.getenv("REMOTE_AGENT_SSE_WRITE_TIMEOUT_S", "10")), + ) + + raise RuntimeError(f"Unsupported AGENT_BACKEND={backend!r}") diff --git a/src/app/livekit/adapters/bridge_gateway.py b/src/app/livekit/adapters/bridge_gateway.py new file mode 100644 index 0000000..c8a6051 --- /dev/null +++ b/src/app/livekit/adapters/bridge_gateway.py @@ -0,0 +1,112 @@ +from __future__ import annotations + +import json +import time +from typing import Any + + +class BridgeGateway: + def __init__( + self, + *, + room: Any, + bridge_identity: str, + protocol: str, + timeline: Any | None = None, + stress_test: bool = False, + ) -> None: + self._room = room + self._bridge_identity = bridge_identity + self._protocol = protocol + self._timeline = timeline + self._stress_test = bool(stress_test) + + async def publish_debug_event(self, event: str, **data: Any) -> None: + """Publish bounded, structured diagnostics to the originating Bridge.""" + if not self._bridge_identity or not self._stress_test: + return + payload = { + "type": "debug_event", + "version": 1, + "source": "agent", + "stress_test": True, + "event": str(event or "unknown").strip(), + "timestamp_ms": round(time.time() * 1000), + "protocol": self._protocol, + "room": self._room.name, + "data": data, + } + await self._room.local_participant.publish_data( + json.dumps(payload, ensure_ascii=False), + reliable=True, + destination_identities=[self._bridge_identity], + topic="agent.debug", + ) + + async def notify_stage_done(self, reason: str = "stage_done") -> None: + if not self._bridge_identity: + return + + payload = { + "type": "stage", + "stage": "DONE", + "reason": reason, + "room": self._room.name, + "protocol": self._protocol, + } + await self._room.local_participant.publish_data( + json.dumps(payload, ensure_ascii=False), + reliable=True, + destination_identities=[self._bridge_identity], + topic="agent.stage", + ) + if self._timeline is not None: + self._timeline.emit( + "bridge_done_notified", + reason=reason, + destination=self._bridge_identity, + ) + + async def notify_stop( + self, + *, + status: str, + reason: str, + resource: str = "", + failed_resources: tuple[str, ...] = (), + phase: str = "in_session", + ) -> None: + if not self._bridge_identity: + return + + payload = { + "type": "stage", + "stage": "DONE", + "status": str(status or "").strip(), + "reason": str(reason or "").strip(), + "room": self._room.name, + "protocol": self._protocol, + } + if resource: + payload["resource"] = str(resource).strip() + normalized_failed_resources = [str(item).strip() for item in failed_resources if str(item).strip()] + if normalized_failed_resources: + payload["failed_resources"] = normalized_failed_resources + if phase: + payload["phase"] = str(phase).strip() + + await self._room.local_participant.publish_data( + json.dumps(payload, ensure_ascii=False), + reliable=True, + destination_identities=[self._bridge_identity], + topic="agent.stage", + ) + if self._timeline is not None: + self._timeline.emit( + "bridge_stop_notified", + status=payload["status"], + reason=payload["reason"], + resource=payload.get("resource", ""), + phase=payload.get("phase", ""), + destination=self._bridge_identity, + ) diff --git a/src/app/livekit/adapters/export_service.py b/src/app/livekit/adapters/export_service.py new file mode 100644 index 0000000..56c855a --- /dev/null +++ b/src/app/livekit/adapters/export_service.py @@ -0,0 +1,11 @@ +from __future__ import annotations + +import asyncio +from typing import Any + +from app.utils.export import json_to_csv + + +class ExportService: + async def export_session(self, output: Any, session_id: str) -> None: + await asyncio.to_thread(json_to_csv, output, session_id) diff --git a/src/app/livekit/adapters/fake_remote_ws_adapter.py b/src/app/livekit/adapters/fake_remote_ws_adapter.py new file mode 100644 index 0000000..0c350ce --- /dev/null +++ b/src/app/livekit/adapters/fake_remote_ws_adapter.py @@ -0,0 +1,428 @@ +from __future__ import annotations + +import asyncio +import os +from typing import Any, Dict, List, Mapping, Optional, Sequence + +from app.livekit.adapters.agent_backend import BackendReply +from app.ws_gateway.fake_remote_agent import build_fake_remote_agent_response + + +class FakeRemoteWSAdapter: + def __init__( + self, + *, + intro: str, + backend_label: str = "remote_ws_fake", + request_context: Optional[Dict[str, Any]] = None, + timeline: Any | None = None, + default_stage: str = "PRESENTATION", + delay_ms: int | None = None, + responses: Sequence[str] | None = None, + ) -> None: + self._intro = intro + self._backend_label = str(backend_label or "remote_ws_fake").strip().lower() + self._request_context = dict(request_context or {}) + self._timeline = timeline + self._default_stage = (default_stage or "PRESENTATION").strip().upper() + configured_delay = ( + os.getenv("FAKE_AGENT_DELAY_MS", "0") if delay_ms is None else delay_ms + ) + try: + self._delay_s = max(0.0, min(180.0, float(configured_delay) / 1000.0)) + except (TypeError, ValueError): + self._delay_s = 0.0 + self._responses = tuple(str(item).strip() for item in (responses or ())) + self._response_index = 0 + self._scripted_final_reply: Optional[BackendReply] = None + + self._elegibility = True + self._protocol = "" + self._last_stage = "INTRO" + self._pending_interrupt: Optional[Dict[str, Any]] = None + self._pending_events: List[Dict[str, Any]] = [] + + self._end_lock = asyncio.Lock() + self._ended = False + self._end_reply: Optional[BackendReply] = None + + @staticmethod + def _normalize_agent_name(value: Any) -> str: + raw = str(value or "").strip().lower() + aliases = { + "conta": "conta", + "contas": "conta", + "ofert": "oferta", + "oferta": "oferta", + "ofertas": "oferta", + "cobra": "cobranca", + "cobranca": "cobranca", + "cobrança": "cobranca", + "cobrancas": "cobranca", + "cobranças": "cobranca", + } + return aliases.get(raw, raw) + + def _current_agent_name(self) -> str: + return self._normalize_agent_name( + self._request_context.get("agent") + or self._request_context.get("Agent") + or self._request_context.get("agente") + or "" + ) + + def _request_field(self, *keys: str) -> str: + for key in keys: + value = self._request_context.get(key) + if value not in (None, ""): + return str(value).strip() + + lowered = {str(key).lower(): value for key, value in self._request_context.items()} + for key in keys: + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return str(value).strip() + + return "" + + def _current_invoice_number(self) -> str: + return self._request_field( + "current_invoice_number", + "currentInvoiceNumber", + "ID_FATURA", + "idFatura", + "id_fatura", + "IdFatura", + ) + + def _current_msisdn(self) -> str: + return self._request_field("msisdn", "GSM", "gsm", "NUM_TELEFONE") + + def _current_channel(self) -> str: + return self._request_field("channel", "Channel", "canal") or "SUPERVISOR" + + @staticmethod + def _extract_user_text(user_input: Any) -> str: + if isinstance(user_input, str): + return user_input.strip() + if isinstance(user_input, Mapping): + for key in ("text", "transcript", "utterance", "message", "content"): + value = user_input.get(key) + if value: + return str(value).strip() + return "" + return str(user_input or "").strip() + + def _base_payload(self) -> Dict[str, Any]: + payload = { + "agent": self._current_agent_name(), + "RouterCallKeyDay": str(self._request_context.get("RouterCallKeyDay") or "").strip(), + "RouterCallKey": str(self._request_context.get("RouterCallKey") or "").strip(), + "ANI": str(self._request_context.get("ANI") or "").strip(), + "GSM": str(self._request_context.get("GSM") or "").strip(), + "callIdGed": str(self._request_context.get("callIdGed") or "").strip(), + "protocol": self._protocol, + "stage": self._last_stage or self._default_stage, + } + id_fatura = str( + self._request_context.get("ID_FATURA") + or self._request_context.get("id_fatura") + or self._request_context.get("IdFatura") + or "" + ).strip() + if id_fatura: + payload["ID_FATURA"] = id_fatura + return payload + + def _build_turn_payload(self, user_input: Any) -> Dict[str, Any]: + text = self._extract_user_text(user_input) + payload = self._base_payload() + if self._current_agent_name() == "conta": + payload = { + "message": text, + "channel": self._current_channel(), + "msisdn": self._current_msisdn(), + } + current_invoice_number = self._current_invoice_number() + if current_invoice_number: + payload["current_invoice_number"] = current_invoice_number + if self._pending_interrupt is not None: + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + return { + "action": "chat", + "payload": payload, + "_agent": self._current_agent_name(), + "_stage": self._last_stage or self._default_stage, + } + + payload["text"] = text + if self._pending_interrupt is not None: + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + return payload + + def _build_end_payload(self) -> Dict[str, Any]: + payload = self._base_payload() + if self._current_agent_name() == "conta": + payload["msisdn"] = self._current_msisdn() + payload["channel"] = self._current_channel() + current_invoice_number = self._current_invoice_number() + if current_invoice_number: + payload["current_invoice_number"] = current_invoice_number + if self._pending_interrupt is not None: + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + return { + "action": "end", + "payload": payload, + "_agent": self._current_agent_name(), + "_stage": self._last_stage or self._default_stage, + } + + payload["type"] = "end" + if self._pending_interrupt is not None: + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + return payload + + def _add_pending_speech_interruption(self, payload: Dict[str, Any]) -> None: + if self._pending_interrupt is None: + return + interruption = dict(self._pending_interrupt) + interruption_key = str( + interruption.pop("_interruption_field", "speech_interruption") + ) + payload[interruption_key] = interruption + + @staticmethod + def _reply_text(response: Mapping[str, Any]) -> str: + result = response.get("result") + if isinstance(result, Mapping): + content = result.get("content") + if content: + return str(content).strip() + text = response.get("text") + if text: + return str(text).strip() + return "" + + @staticmethod + def _export_payload(response: Mapping[str, Any]) -> Any: + return response.get("result") + + def _to_backend_reply(self, response: Mapping[str, Any]) -> BackendReply: + stage = str(response.get("stage") or self._last_stage or self._default_stage).strip().upper() + done = stage == "DONE" or str(response.get("type") or "").strip().lower() == "done" + return BackendReply( + stage=stage, + text=self._reply_text(response), + done=done, + export_payload=self._export_payload(response), + ) + + def _clear_pending_state(self) -> None: + self._pending_interrupt = None + self._pending_events.clear() + + async def prepare( + self, + elegibility: bool, + protocol: str, + ) -> None: + if self._timeline is not None: + self._timeline.emit( + "backend_prepare_started", + backend=self._backend_label, + backend_family="remote_ws_fake", + agent=self._current_agent_name(), + ) + self._elegibility = bool(elegibility) + self._protocol = str(protocol or "") + self._last_stage = "INTRO" + self._pending_interrupt = None + self._pending_events.clear() + self._ended = False + self._end_reply = None + self._response_index = 0 + self._scripted_final_reply = None + if self._timeline is not None: + self._timeline.emit( + "backend_prepare_completed", + backend=self._backend_label, + backend_family="remote_ws_fake", + agent=self._current_agent_name(), + protocol=self._protocol, + ) + + async def run(self, user_input: Any) -> BackendReply: + if self._scripted_final_reply is not None: + return self._scripted_final_reply + + payload = self._build_turn_payload(user_input) + if self._timeline is not None: + self._timeline.emit( + "remote_agent_request", + backend=self._backend_label, + backend_family="remote_ws_fake", + agent=self._current_agent_name(), + action=str(payload.get("action") or "chat"), + payload=payload, + ) + + if self._delay_s: + await asyncio.sleep(self._delay_s) + if self._responses: + index = self._response_index + text = self._responses[index] + if index == len(self._responses) - 1: + stage = "DONE" + elif index == len(self._responses) - 2: + stage = "FORMALIZATION" + else: + stage = "ARGUMENTATION" + done = stage == "DONE" + reply = BackendReply( + stage=stage, + text=text, + done=done, + export_payload=( + { + "type": "final", + "content": text, + "tool_calls": [], + "result": [{"status": "ok", "reason": "fake_done"}], + } + if done + else None + ), + ) + self._response_index += 1 + if done: + self._scripted_final_reply = reply + else: + response = build_fake_remote_agent_response(payload) + reply = self._to_backend_reply(response) + self._last_stage = reply.stage or self._last_stage or self._default_stage + self._clear_pending_state() + + if self._timeline is not None: + self._timeline.emit( + "remote_agent_response", + backend=self._backend_label, + backend_family="remote_ws_fake", + agent=self._current_agent_name(), + stage=reply.stage, + done=reply.done, + text_len=len(reply.text), + has_result=reply.export_payload is not None, + ) + return reply + + async def set_interruption( + self, + interrupted: bool, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + if not interrupted: + self._pending_interrupt = None + else: + self._pending_interrupt = { + "speech_id": str(speech_id or "").strip(), + "heard_text": listened_text or "", + } + if self._timeline is not None: + self._timeline.emit( + "backend_set_interruption", + backend=self._backend_label, + backend_family="remote_ws_fake", + agent=self._current_agent_name(), + interrupted=bool(interrupted), + skipped=bool(skipped), + listened_text=listened_text, + speech_id=str(speech_id or "").strip(), + ) + + async def set_processing_interruption( + self, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + await self.set_interruption( + True, + listened_text=listened_text, + skipped=skipped, + speech_id=speech_id, + ) + if self._pending_interrupt is not None: + self._pending_interrupt["_interruption_field"] = "processing_interruption" + + async def inject_idle_nudge(self, nudge_text: str) -> None: + text = str(nudge_text or "").strip() + if not text: + return + event = {"type": "idle_nudge", "text": text} + # Cada frase de inatividade substitui a anterior: o cliente responde ao + # que ouviu por ultimo, e as intermediarias so empilham falas do agente + # no historico remoto -- inclusive o aviso de encerramento, que passa a + # constar como dito logo antes de a conversa seguir normalmente. + if self._pending_events and self._pending_events[-1].get("type") == "idle_nudge": + self._pending_events[-1] = event + else: + self._pending_events.append(event) + if self._timeline is not None: + self._timeline.emit( + "backend_idle_nudge_buffered", + backend=self._backend_label, + backend_family="remote_ws_fake", + agent=self._current_agent_name(), + text=text, + ) + + def supports_server_push(self) -> bool: + return False + + async def wait_for_server_push(self) -> BackendReply | None: + return None + + async def end_service_once(self) -> BackendReply: + async with self._end_lock: + if self._ended: + return self._end_reply or BackendReply(stage="DONE", done=True, export_payload=[]) + + self._ended = True + if self._timeline is not None: + self._timeline.emit( + "backend_end_started", + backend=self._backend_label, + backend_family="remote_ws_fake", + agent=self._current_agent_name(), + ) + response = build_fake_remote_agent_response(self._build_end_payload()) + reply = self._to_backend_reply(response) + if not reply.done: + reply = BackendReply( + stage="DONE", + text=reply.text, + done=True, + export_payload=reply.export_payload, + ) + self._last_stage = reply.stage + self._clear_pending_state() + self._end_reply = reply + if self._timeline is not None: + self._timeline.emit( + "backend_end_completed", + backend=self._backend_label, + backend_family="remote_ws_fake", + agent=self._current_agent_name(), + result_type=type(reply.export_payload).__name__, + ) + return reply diff --git a/src/app/livekit/adapters/fake_tts.py b/src/app/livekit/adapters/fake_tts.py new file mode 100644 index 0000000..b4271eb --- /dev/null +++ b/src/app/livekit/adapters/fake_tts.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +import asyncio + +from livekit.agents import tts, utils +from livekit.agents.types import DEFAULT_API_CONNECT_OPTIONS, APIConnectOptions + +from app.providers.tts import CHANNELS, SAMPLE_RATE +from app.providers.tts import FakeTTS as FakeProviderTTS + + +class FakeTTS(tts.TTS): + def __init__(self) -> None: + super().__init__( + capabilities=tts.TTSCapabilities(streaming=False, aligned_transcript=False), + sample_rate=SAMPLE_RATE, + num_channels=CHANNELS, + ) + self._provider_client = FakeProviderTTS() + + @property + def model(self) -> str: + return "fake-tone" + + @property + def provider(self) -> str: + return "fake" + + def synthesize( + self, + text: str, + *, + conn_options: APIConnectOptions = DEFAULT_API_CONNECT_OPTIONS, + ) -> ChunkedStream: + return ChunkedStream(tts=self, input_text=text, conn_options=conn_options) + + async def aclose(self) -> None: + return None + + +class ChunkedStream(tts.ChunkedStream): + def __init__( + self, + *, + tts: FakeTTS, + input_text: str, + conn_options: APIConnectOptions, + ) -> None: + super().__init__(tts=tts, input_text=input_text, conn_options=conn_options) + self._tts: FakeTTS = tts + + async def _run(self, output_emitter: tts.AudioEmitter) -> None: + output_emitter.initialize( + request_id=utils.shortuuid(), + sample_rate=self._tts.sample_rate, + num_channels=self._tts.num_channels, + mime_type="audio/pcm", + ) + + audio = await asyncio.to_thread( + self._tts._provider_client.synthesize_pcm16k, + self._input_text, + ) + output_emitter.push(audio) + output_emitter.flush() diff --git a/src/app/livekit/adapters/pipeline_adapter.py b/src/app/livekit/adapters/pipeline_adapter.py new file mode 100644 index 0000000..6788245 --- /dev/null +++ b/src/app/livekit/adapters/pipeline_adapter.py @@ -0,0 +1,202 @@ +from __future__ import annotations + +import asyncio +from typing import Any, Dict, Optional + +from agent.pipeline.customer_pipeline_langgraph import CustomerPipeline +from app.livekit.adapters.agent_backend import BackendReply +from app.utils.logging import setup_minimal_logging + +logger = setup_minimal_logging() + + +class PipelineAdapter: + def __init__( + self, + *, + session_data: Dict[str, Any], + intro: str, + timeline: Any | None = None, + streaming: bool = False, + ) -> None: + self._session_data = session_data + self._intro = intro + self._timeline = timeline + self._streaming = streaming + + self._pipeline: Optional[CustomerPipeline] = None + self._last_stage = "INTRO" + self._end_lock = asyncio.Lock() + self._ended = False + self._end_reply: Optional[BackendReply] = None + + def _require_pipeline(self) -> CustomerPipeline: + if self._pipeline is None: + raise RuntimeError("PipelineAdapter used before prepare() completed") + return self._pipeline + + def _call_set_interruption( + self, + pipeline: CustomerPipeline, + interrupted: bool, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + fn = getattr(pipeline, "set_interruption", None) + if not callable(fn): + return + + try: + fn(interrupted, listened_text, skipped, speech_id) + return + except TypeError: + pass + + try: + fn(interrupted, listened_text, skipped) + return + except TypeError: + pass + + try: + fn(interrupted, listened_text) + return + except TypeError: + pass + + fn(interrupted) + + async def prepare( + self, + elegibility: bool, + protocol: str, + ) -> None: + if self._timeline is not None: + self._timeline.emit( + "backend_prepare_started", + backend="langgraph", + elegibility=bool(elegibility), + ) + + def _build_pipeline() -> CustomerPipeline: + pipeline = CustomerPipeline(self._session_data, streaming=self._streaming) + pipeline.intro = self._intro + return pipeline + + pipeline = await asyncio.to_thread(_build_pipeline) + self._pipeline = pipeline + + def _prepare() -> None: + pipeline.prepare(elegibility, protocol) + + await asyncio.to_thread(_prepare) + if self._timeline is not None: + self._timeline.emit("backend_prepare_completed", backend="langgraph") + + async def run(self, user_input: Any) -> BackendReply: + pipeline = self._require_pipeline() + result = await asyncio.to_thread(pipeline.run, user_input) + if not isinstance(result, tuple) or len(result) != 2: + raise RuntimeError(f"Unexpected pipeline.run() result: {type(result)!r}") + + stage_raw, output_raw = result + stage = str(stage_raw or self._last_stage or "PRESENTATION").strip().upper() + text = str(output_raw or "").strip() + self._last_stage = stage + reply = BackendReply( + stage=stage, + text=text, + done=stage == "DONE", + export_payload=None, + ) + if self._timeline is not None: + self._timeline.emit( + "backend_run_completed", + backend="langgraph", + stage=reply.stage, + output_len=len(reply.text), + ) + return reply + + async def set_interruption( + self, + interrupted: bool, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + pipeline = self._require_pipeline() + if self._timeline is not None: + self._timeline.emit( + "backend_set_interruption", + backend="langgraph", + interrupted=bool(interrupted), + skipped=bool(skipped), + listened_text=listened_text, + speech_id=speech_id, + ) + self._call_set_interruption(pipeline, interrupted, listened_text, skipped, speech_id) + + async def set_processing_interruption( + self, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + await self.set_interruption( + True, + listened_text=listened_text, + skipped=skipped, + speech_id=speech_id, + ) + + async def inject_idle_nudge(self, nudge_text: str) -> None: + pipeline = self._require_pipeline() + + inject_user = getattr(getattr(pipeline, "agent", None), "inject_user_message", None) + inject_ai = getattr(getattr(pipeline, "agent", None), "inject_ai_message", None) + if callable(inject_user) and callable(inject_ai): + inject_user("") + inject_ai(nudge_text) + + self._call_set_interruption(pipeline, False, "", False, "") + + update_auto = getattr(pipeline, "update_langfuse_auto", None) + if callable(update_auto): + update_auto("###idle###", nudge_text) + + async def end_service_once(self) -> BackendReply: + async with self._end_lock: + if self._ended: + return self._end_reply or BackendReply(stage="DONE", done=True, export_payload=[]) + + self._ended = True + if self._timeline is not None: + self._timeline.emit("backend_end_started", backend="langgraph") + pipeline = self._pipeline + if pipeline is None: + self._end_reply = BackendReply(stage="DONE", done=True, export_payload=[]) + return self._end_reply + + end_output: Any + try: + end_output = await pipeline.end_service() + except Exception: + logger.exception("[pipeline] end_service() falhou") + end_output = [] + + self._end_reply = BackendReply( + stage="DONE", + text="", + done=True, + export_payload=end_output if end_output is not None else [], + ) + + if self._timeline is not None: + self._timeline.emit( + "backend_end_completed", + backend="langgraph", + result_type=type(self._end_reply.export_payload).__name__, + ) + return self._end_reply diff --git a/src/app/livekit/adapters/remote_agent_sse_adapter.py b/src/app/livekit/adapters/remote_agent_sse_adapter.py new file mode 100644 index 0000000..dd93c99 --- /dev/null +++ b/src/app/livekit/adapters/remote_agent_sse_adapter.py @@ -0,0 +1,1817 @@ +from __future__ import annotations + +import asyncio +import json +import os +import re +import time +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import Any, Dict, List, Mapping, Optional + +import httpx + +from app.livekit.adapters.agent_backend import BackendReply +from app.livekit.policies.agent_finalization import final_stop_from_agent_result +from app.utils.logging import setup_minimal_logging +from app.utils.turn_ids import next_turn_message_id + +logger = setup_minimal_logging() + +_STREAM_TYPES = {"chunk", "delta", "partial", "token", "progress"} +_TERMINAL_TYPES = {"done", "final", "message", "output", "response", "result", "proactive_result"} +_PREFETCH_TERMINAL_TYPES = {"prefetch_done", "prefetch_skipped", "prefetch_failed"} +_TEXT_KEYS = ("text", "response", "output", "message", "content") +_USER_TEXT_KEYS = ("text", "transcript", "utterance", "message", "content") +_OFERTA_INITIAL_MESSAGE_ENV = "REMOTE_AGENT_OFERTA_INITIAL_MESSAGE" +_OFERTA_INITIAL_MESSAGE_DEFAULT = "inicio_atendimento" +_OFERTA_SERVICE_STATUS_RESULT_TYPE = { + "RESOLVED": "resolvido", +} +_OFERTA_NON_FINAL_SERVICE_STATUSES = {"UNRESOLVED", "RESOLVED_WITH_NEW_REQUEST"} +_FALSE_ENV_VALUES = {"0", "false", "no", "off", "disable", "disabled"} +_TRUE_ENV_VALUES = {"1", "true", "yes", "on", "enable", "enabled"} +_CONTA_USER_RESPONSE_MESSAGE_TYPES = {"ready", "final"} + + +@dataclass(slots=True) +class _RemoteAgentReply: + stage: str + text: str + done: bool + result: Any + event_type: str = "" + metadata: Any = None + + +class RemoteAgentSSEAdapter: + def __init__( + self, + *, + intro: str, + url: str, + backend_label: str = "remote_sse", + url_by_agent: Optional[Dict[str, str]] = None, + request_context: Optional[Dict[str, Any]] = None, + timeline: Any | None = None, + streaming: bool = False, + default_stage: str = "PRESENTATION", + connect_timeout_s: float = 10.0, + read_timeout_s: float = 45.0, + write_timeout_s: float = 10.0, + ) -> None: + self._intro = intro + self._url = url + self._backend_label = str(backend_label or "remote_sse").strip().lower() + self._url_by_agent = { + self._normalize_agent_name(key): value.strip() + for key, value in (url_by_agent or {}).items() + if value and value.strip() + } + self._request_context = dict(request_context or {}) + self._timeline = timeline + self._streaming = streaming + self._default_stage = (default_stage or "PRESENTATION").strip().upper() + self._connect_timeout_s = connect_timeout_s + self._read_timeout_s = read_timeout_s + self._write_timeout_s = write_timeout_s + + self._elegibility = True + self._remote_session_id = "" + self._last_stage = "INTRO" + self._pending_interrupt: Optional[Dict[str, Any]] = None + self._pending_events: List[Dict[str, Any]] = [] + + self._end_lock = asyncio.Lock() + self._ended = False + self._end_reply: Optional[BackendReply] = None + self._session_lock = asyncio.Lock() + self._pending_ready_reply: Optional[_RemoteAgentReply] = None + self._push_queue: asyncio.Queue[_RemoteAgentReply] = asyncio.Queue() + self._session_actions: set[str] | None = None + self._protocol = "" + self._turn_seq = 0 + self._offer_stream_task: asyncio.Task[None] | None = None + self._conta_feedback_streams = 0 + self._conta_feedback_event = asyncio.Event() + self._active_request_message_id = "" + + @staticmethod + def _normalize_agent_name(value: Any) -> str: + raw = str(value or "").strip().lower() + aliases = { + "conta": "conta", + "contas": "conta", + "ofert": "oferta", + "oferta": "oferta", + "ofertas": "oferta", + "cobra": "cobranca", + "cobranca": "cobranca", + "cobrança": "cobranca", + "cobrancas": "cobranca", + "cobranças": "cobranca", + } + return aliases.get(raw, raw) + + @staticmethod + def _string_or_empty(value: Any) -> str: + return str(value or "").strip() + + @classmethod + def _message_id_from_mapping(cls, payload: Mapping[str, Any] | None) -> str: + if not isinstance(payload, Mapping): + return "" + + for key in ("message_id", "messageId"): + value = payload.get(key) + if value not in (None, ""): + return cls._string_or_empty(value) + + for container_key in ("metadata", "data", "payload", "context"): + metadata = payload.get(container_key) + if not isinstance(metadata, Mapping): + continue + for key in ("message_id", "messageId"): + value = metadata.get(key) + if value not in (None, ""): + return cls._string_or_empty(value) + + return "" + + @classmethod + def _metadata_with_message_id(cls, metadata: Any, message_id: Any) -> Dict[str, Any] | None: + normalized_message_id = cls._string_or_empty(message_id) + if not normalized_message_id: + return dict(metadata) if isinstance(metadata, Mapping) else metadata + + merged: Dict[str, Any] = dict(metadata) if isinstance(metadata, Mapping) else {} + merged.setdefault("message_id", normalized_message_id) + return merged + + @classmethod + def _reply_with_message_id( + cls, + reply: _RemoteAgentReply, + message_id: Any, + ) -> _RemoteAgentReply: + metadata = cls._metadata_with_message_id(reply.metadata, message_id) + if metadata is not reply.metadata: + reply.metadata = metadata + return reply + + @staticmethod + def _decode_message(raw_message: Any) -> Any: + if isinstance(raw_message, bytes): + raw_message = raw_message.decode("utf-8") + if not isinstance(raw_message, str): + raise RuntimeError(f"Unsupported SSE payload type: {type(raw_message)!r}") + + message = raw_message.strip() + if not message: + return {} + + try: + return json.loads(message) + except json.JSONDecodeError: + return message + + def _extract_text_from_value(self, value: Any) -> str: + if isinstance(value, str): + return value.strip() + if isinstance(value, Mapping): + for key in ("result", *_TEXT_KEYS): + text = self._extract_text_from_value(value.get(key)) + if text: + return text + return "" + if isinstance(value, list): + for item in value: + text = self._extract_text_from_value(item) + if text: + return text + return "" + + def _extract_user_text(self, user_input: Any) -> str: + if isinstance(user_input, str): + return user_input.strip() + + if isinstance(user_input, Mapping): + for key in _USER_TEXT_KEYS: + text = self._extract_text_from_value(user_input.get(key)) + if text: + return text + + for value in user_input.values(): + text = self._extract_text_from_value(value) + if text: + return text + + if isinstance(user_input, list): + for item in user_input: + text = self._extract_user_text(item) + if text: + return text + + return str(user_input).strip() + + def _current_timestamp(self) -> str: + return datetime.now(timezone.utc).isoformat() + + def _request_field(self, *keys: str) -> str: + for key in keys: + value = self._request_context.get(key) + if value not in (None, ""): + return self._string_or_empty(value) + + lowered = {str(key).lower(): value for key, value in self._request_context.items()} + for key in keys: + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return self._string_or_empty(value) + + return "" + + def _current_agent_name(self) -> str: + return self._normalize_agent_name( + self._request_field("agent", "Agent", "agente") + ) + + def _is_oferta_agent(self) -> bool: + return self._current_agent_name() == "oferta" + + def _resolve_url(self) -> str: + agent_name = self._current_agent_name() + if agent_name: + routed_url = self._url_by_agent.get(agent_name, "") + if routed_url: + return routed_url + return self._url + + def _current_invoice_number(self) -> str: + return self._request_field( + "invoice_id", + "invoiceId", + "current_invoice_number", + "currentInvoiceNumber", + "ID_FATURA", + "idFatura", + "id_fatura", + "IdFatura", + ) + + def _current_msisdn(self) -> str: + return self._request_field("msisdn", "GSM", "gsm", "NUM_TELEFONE") + + def _current_ani(self) -> str: + return self._request_field("ANI", "ani") + + def _current_session_id(self) -> str: + return self._request_field("session_id", "sessionId") + + def _current_channel_id(self) -> str: + channel_id = self._request_field("channelId", "channel_id") + if channel_id: + return channel_id + return "ura" + + def _current_ura_call_id(self) -> str: + return self._request_field("uraCallId", "ura_call_id", "callIdGed") + + def _current_protocol_number(self) -> str: + return self._request_field("protocol_id", "protocolId", "protocolo", "protocolNumber", "protocol") + + def _message_id_from_user_input(self, user_input: Any) -> str: + if not isinstance(user_input, Mapping): + return "" + + for key in ("message_id", "messageId"): + value = user_input.get(key) + if value not in (None, ""): + return self._string_or_empty(value) + + for container_key in ("metadata", "data", "payload"): + nested = user_input.get(container_key) + if not isinstance(nested, Mapping): + continue + for key in ("message_id", "messageId"): + value = nested.get(key) + if value not in (None, ""): + return self._string_or_empty(value) + + return "" + + def _query_message_id(self, user_input: Any = None) -> str: + return ( + self._message_id_from_user_input(user_input) + or self._request_field("message_id", "messageId") + or self._next_message_id() + ) + + def _current_channel(self) -> str: + channel = self._request_field("channel", "Channel", "canal") + if channel: + return channel + + configured = self._string_or_empty( + os.getenv("REMOTE_AGENT_CHANNEL") + or os.getenv("REMOTE_AGENT_DEFAULT_CHANNEL") + ) + return configured or "SUPERVISOR" + + def _stream_open_params(self) -> Dict[str, str]: + params: Dict[str, str] = {} + + msisdn = self._current_msisdn() + if msisdn: + params["msisdn"] = msisdn + + invoice_number = self._current_invoice_number() + if invoice_number: + params["invoice_id"] = invoice_number + + if self._current_agent_name() != "conta": + return params + + ani = self._current_ani() + if ani: + params["ani"] = ani + + protocol_id = self._current_protocol_number() + if protocol_id: + params["protocol_id"] = protocol_id + + session_id = self._current_session_id() + if session_id: + params["session_id"] = session_id + + message_id = self._query_message_id() + if message_id: + params["message_id"] = message_id + + channel_id = self._current_channel_id() + if channel_id: + params["channelId"] = channel_id + + ura_call_id = self._current_ura_call_id() + if ura_call_id: + params["uraCallId"] = ura_call_id + + return params + + def _action_params( + self, + user_input: Any = None, + *, + message_id: str = "", + ) -> Dict[str, str]: + params: Dict[str, str] = {} + if self._remote_session_id: + params["session_id"] = self._remote_session_id + + if self._current_agent_name() != "conta": + return params + + if "session_id" not in params: + session_id = self._current_session_id() + if session_id: + params["session_id"] = session_id + + ani = self._current_ani() + if ani: + params["ani"] = ani + + protocol_id = self._current_protocol_number() + if protocol_id: + params["protocol_id"] = protocol_id + + request_message_id = message_id or self._query_message_id(user_input) + if request_message_id: + params["message_id"] = request_message_id + + channel_id = self._current_channel_id() + if channel_id: + params["channelId"] = channel_id + + ura_call_id = self._current_ura_call_id() + if ura_call_id: + params["uraCallId"] = ura_call_id + + return params + + def _stream_headers(self) -> Dict[str, str]: + return { + "Accept": "text/event-stream", + "Cache-Control": "no-cache", + } + + def _oferta_headers(self) -> Dict[str, str]: + headers = self._stream_headers() + headers["Content-Type"] = "application/json" + headers["Channel-id"] = self._current_channel_id() or "ura" + return headers + + def _client_timeout(self) -> httpx.Timeout: + return httpx.Timeout( + connect=self._connect_timeout_s, + read=self._read_timeout_s, + write=self._write_timeout_s, + pool=self._connect_timeout_s, + ) + + def _update_session_actions(self, payload: Mapping[str, Any]) -> None: + if "actions" not in payload: + return + + raw_actions = payload.get("actions") + parsed_actions: set[str] = set() + if isinstance(raw_actions, list): + for item in raw_actions: + action = str(item or "").strip().lower() + if action: + parsed_actions.add(action) + elif isinstance(raw_actions, str): + action = raw_actions.strip().lower() + if action: + parsed_actions.add(action) + + self._session_actions = parsed_actions + + def _update_session_id(self, payload: Mapping[str, Any]) -> None: + for key in ("session_id", "sessionId"): + value = payload.get(key) + if value not in (None, ""): + self._remote_session_id = self._string_or_empty(value) + return + + def _session_supports_action(self, action: str) -> bool: + normalized_action = str(action or "").strip().lower() + if not normalized_action: + return False + + if self._session_actions is None: + return normalized_action != "end" + + return normalized_action in self._session_actions + + @staticmethod + def _is_unsupported_action_error(error: BaseException, action: str) -> bool: + action_name = str(action or "").strip().lower() + if not action_name: + return False + + message = str(error or "").strip().lower() + return "não suportada" in message and action_name in message + + def _emit_request( + self, + *, + method: str, + url: str, + action: str, + params: Mapping[str, Any], + payload: Dict[str, Any] | None = None, + ) -> None: + if self._timeline is not None: + self._timeline.emit( + "remote_agent_request", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + method=str(method or "").upper(), + url=url, + action=action, + params=dict(params), + payload=payload, + ) + + def _clear_push_queue(self) -> None: + while True: + try: + self._push_queue.get_nowait() + except asyncio.QueueEmpty: + return + + def supports_server_push(self) -> bool: + if not self._push_queue.empty(): + return True + if not self._is_oferta_agent() or self._offer_stream_task is None: + return False + if not self._offer_stream_task.done(): + return True + try: + return self._offer_stream_task.exception() is not None + except asyncio.CancelledError: + return False + + def supports_inflight_backend_push(self) -> bool: + return self._current_agent_name() == "conta" + + async def wait_for_server_push(self) -> BackendReply | None: + reply = await self._next_pushed_reply() + if reply is None: + return None + + self._last_stage = reply.stage or self._last_stage or self._default_stage + return BackendReply( + stage=self._last_stage, + text=reply.text, + done=bool(reply.done) or self._last_stage == "DONE", + export_payload=reply.result, + metadata=reply.metadata, + ) + + async def _next_pushed_reply(self) -> _RemoteAgentReply | None: + try: + return self._push_queue.get_nowait() + except asyncio.QueueEmpty: + pass + + if self._current_agent_name() == "conta": + return await self._next_conta_feedback_reply() + + if not self._is_oferta_agent(): + return None + + stream_task = self._offer_stream_task + if stream_task is None: + return None + + if stream_task.done(): + return self._reply_from_finished_offer_stream(stream_task) + + get_task = asyncio.create_task(self._push_queue.get(), name="remote_sse_oferta_push_get") + try: + done, _pending = await asyncio.wait( + {get_task, stream_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + except asyncio.CancelledError: + get_task.cancel() + await asyncio.gather(get_task, return_exceptions=True) + raise + + if get_task in done: + return get_task.result() + + get_task.cancel() + await asyncio.gather(get_task, return_exceptions=True) + try: + return self._push_queue.get_nowait() + except asyncio.QueueEmpty: + return self._reply_from_finished_offer_stream(stream_task) + + def _reply_from_finished_offer_stream(self, stream_task: asyncio.Task[None]) -> _RemoteAgentReply | None: + try: + exc = stream_task.exception() + except asyncio.CancelledError: + return None + if exc is not None: + raise RuntimeError("Remote oferta SSE stream failed") from exc + return None + + async def _next_conta_feedback_reply(self) -> _RemoteAgentReply | None: + saw_stream = False + while True: + try: + return self._push_queue.get_nowait() + except asyncio.QueueEmpty: + pass + + if self._conta_feedback_streams <= 0: + if saw_stream: + return None + self._conta_feedback_event.clear() + try: + await asyncio.wait_for(self._conta_feedback_event.wait(), timeout=1.0) + except asyncio.TimeoutError: + return None + continue + + saw_stream = True + get_task = asyncio.create_task(self._push_queue.get(), name="remote_sse_conta_feedback_get") + state_task = asyncio.create_task( + self._conta_feedback_event.wait(), + name="remote_sse_conta_feedback_state", + ) + try: + done, _pending = await asyncio.wait( + {get_task, state_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + except asyncio.CancelledError: + get_task.cancel() + state_task.cancel() + await asyncio.gather(get_task, state_task, return_exceptions=True) + raise + + if get_task in done: + state_task.cancel() + await asyncio.gather(state_task, return_exceptions=True) + return get_task.result() + + get_task.cancel() + await asyncio.gather(get_task, return_exceptions=True) + self._conta_feedback_event.clear() + + def _base_payload(self) -> Dict[str, Any]: + payload = { + "timestamp": self._current_timestamp(), + "agent": self._current_agent_name(), + "RouterCallKeyDay": self._request_field("RouterCallKeyDay", "router_call_key_day", "routerCallKeyDay"), + "RouterCallKey": self._request_field("RouterCallKey", "router_call_key", "routerCallKey"), + "ANI": self._request_field("ANI", "ani"), + "GSM": self._request_field("GSM", "gsm", "NUM_TELEFONE"), + "callIdGed": self._request_field("callIdGed"), + } + + invoice_number = self._current_invoice_number() + if invoice_number: + payload["invoice_id"] = invoice_number + if self._current_agent_name() == "conta": + payload["ID_FATURA"] = invoice_number + + return payload + + def _next_message_id(self) -> str: + source = dict(self._request_context) + if self._protocol: + source.setdefault("protocol", self._protocol) + return next_turn_message_id(source) + + def _current_oferta_initial_message(self) -> str: + return ( + self._string_or_empty(os.getenv(_OFERTA_INITIAL_MESSAGE_ENV)) + or _OFERTA_INITIAL_MESSAGE_DEFAULT + ) + + def _env_first(self, *names: str) -> str: + for name in names: + value = self._string_or_empty(os.getenv(name)) + if value: + return value + return "" + + def _verify_tls(self) -> bool: + agent_name = re.sub(r"[^A-Z0-9]+", "_", self._current_agent_name().upper()).strip("_") + candidate_names = [] + if agent_name: + candidate_names.extend( + [ + f"REMOTE_AGENT_SSE_TLS_VERIFY_{agent_name}", + f"REMOTE_AGENT_TLS_VERIFY_{agent_name}", + f"REMOTE_AGENT_SSE_VERIFY_TLS_{agent_name}", + f"REMOTE_AGENT_VERIFY_TLS_{agent_name}", + ] + ) + candidate_names.extend( + [ + "REMOTE_AGENT_SSE_TLS_VERIFY", + "REMOTE_AGENT_TLS_VERIFY", + "REMOTE_AGENT_SSE_VERIFY_TLS", + "REMOTE_AGENT_VERIFY_TLS", + ] + ) + + raw_value = self._env_first(*candidate_names) + if not raw_value: + return True + + normalized = raw_value.strip().lower() + if normalized in _FALSE_ENV_VALUES: + return False + if normalized in _TRUE_ENV_VALUES: + return True + return True + + def _build_oferta_context(self) -> Dict[str, Any]: + context: Dict[str, Any] = {} + + explicit_session_id = self._request_field("session_id", "sessionId") + if explicit_session_id: + context["sessionId"] = explicit_session_id + + protocol_number = self._current_protocol_number() + if protocol_number: + context["protocolNumber"] = protocol_number + context["protocolo"] = protocol_number + + gsm = self._current_msisdn() + if gsm: + context["gsm"] = gsm + + call_id_ged = self._request_field("callIdGed") + if call_id_ged: + context["uraId"] = call_id_ged + context["callIdGed"] = call_id_ged + + optional_fields = ( + ("ani", ("ANI", "ani")), + ("routerCallKey", ("RouterCallKey", "routerCallKey")), + ("routerCallKeyDay", ("RouterCallKeyDay", "routerCallKeyDay")), + ("agent", ("agent",)), + ("assetId", ("assetId", "asset_id")), + ) + for output_key, input_keys in optional_fields: + value = self._request_field(*input_keys) + if value: + context[output_key] = value + + self._add_pending_speech_interruption(context) + if self._pending_events: + context["events"] = list(self._pending_events) + + return context + + def _build_oferta_execute_payload(self, user_input: Any) -> Dict[str, Any]: + user_text = self._extract_user_text(user_input) + message = user_text or self._current_oferta_initial_message() + message_id = self._message_id_from_user_input(user_input) or self._next_message_id() + return { + "messageId": message_id, + "message": message, + "context": self._build_oferta_context(), + } + + def _build_conta_payload( + self, + user_input: Any, + *, + message_id: str = "", + ) -> Dict[str, Any]: + payload = { + "message": self._extract_user_text(user_input), + } + if message_id: + payload["message_id"] = message_id + + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + + return payload + + def _build_turn_payload( + self, + user_input: Any, + *, + message_id: str = "", + ) -> Dict[str, Any]: + if self._current_agent_name() == "conta": + request_message_id = message_id or self._query_message_id(user_input) + return { + "action": "chat", + "payload": self._build_conta_payload( + user_input, + message_id=request_message_id, + ), + } + + payload = self._base_payload() + payload["message"] = self._extract_user_text(user_input) + payload["msisdn"] = self._current_msisdn() + payload["stage"] = self._last_stage + + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + + return { + "action": "chat", + "payload": payload, + } + + def _build_end_payload(self, *, message_id: str = "") -> Dict[str, Any]: + if self._current_agent_name() == "conta": + request_message_id = message_id or self._query_message_id() + return { + "action": "end", + "payload": self._build_conta_payload("", message_id=request_message_id), + } + + payload = self._base_payload() + payload["msisdn"] = self._current_msisdn() + payload["stage"] = self._last_stage + + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + + return { + "action": "end", + "payload": payload, + } + + def _clear_pending_state(self) -> None: + self._pending_interrupt = None + self._pending_events.clear() + + def _extract_stage(self, payload: Mapping[str, Any]) -> str: + stage = payload.get("stage") + if isinstance(stage, str) and stage.strip(): + return stage.strip().upper() + return "" + + def _extract_delta(self, payload: Mapping[str, Any]) -> str: + for key in ("delta", "chunk", "token"): + value = payload.get(key) + if isinstance(value, str) and value.strip(): + return value + return "" + + def _extract_reply_text(self, payload: Mapping[str, Any]) -> str: + message_type = str(payload.get("type") or "").strip().lower() + if message_type == "final": + text = self._extract_text_from_value(payload.get("content")) + if text: + return text + + for key in ("result", *_TEXT_KEYS): + value = payload.get(key) + text = self._extract_text_from_value(value) + if text: + return text + return "" + + @staticmethod + def _bool_or_none(value: Any) -> bool | None: + if isinstance(value, bool): + return value + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"1", "true", "yes", "on"}: + return True + if normalized in {"0", "false", "no", "off"}: + return False + return None + + def _reply_metadata(self, payload: Mapping[str, Any], result: Any = None) -> Dict[str, Any] | None: + raw_metadata = payload.get("metadata") + metadata: Dict[str, Any] = dict(raw_metadata) if isinstance(raw_metadata, Mapping) else {} + + sources: list[Mapping[str, Any]] = [] + if isinstance(result, Mapping): + sources.append(result) + sources.append(payload) + + for source in sources: + message_id = self._message_id_from_mapping(source) + if message_id: + metadata["message_id"] = message_id + break + + for source in sources: + speech_id = self._string_or_empty(source.get("speech_id")) + if speech_id: + metadata["speech_id"] = speech_id + break + + for source in sources: + interruptible = self._bool_or_none(source.get("is_interruptible")) + if interruptible is not None: + metadata["is_interruptible"] = interruptible + break + + if self._current_agent_name() == "conta": + message_type = str(payload.get("type") or "").strip().lower() + result_type = "" + if isinstance(result, Mapping): + result_type = str(result.get("type") or "").strip().lower() + + expects_user_response = ( + result_type == "final" + or (not result_type and message_type in _CONTA_USER_RESPONSE_MESSAGE_TYPES) + ) + terminal_result = final_stop_from_agent_result(result) is not None + protected_speech = ( + message_type in {"ready", "feedback"} + or terminal_result + or (bool(result_type) and result_type != "final") + ) + + if message_type: + metadata["agent_message_type"] = message_type + if result_type: + metadata["agent_result_type"] = result_type + metadata["expects_user_response"] = expects_user_response + metadata["drop_user_input_while_speaking"] = protected_speech + if protected_speech: + metadata["is_interruptible"] = False + + return metadata or None + + def _is_conta_feedback(self, message_type: str) -> bool: + return self._current_agent_name() == "conta" and message_type == "feedback" + + def _is_conta_result_feedback(self, result: Any) -> bool: + return ( + self._current_agent_name() == "conta" + and isinstance(result, Mapping) + and str(result.get("type") or "").strip().lower() == "feedback" + ) + + def _feedback_metadata( + self, + payload: Mapping[str, Any], + base_metadata: Any = None, + ) -> Dict[str, Any]: + metadata: Dict[str, Any] = {} + if isinstance(base_metadata, Mapping): + metadata.update(base_metadata) + metadata["event"] = "feedback" + metadata["payload"] = dict(payload) + return metadata + + def _queue_feedback_reply( + self, + *, + payload: Mapping[str, Any], + stage: str, + text: str, + metadata: Any = None, + ) -> None: + feedback_text = (text or "").strip() + if not feedback_text: + return + + metadata = self._metadata_with_message_id(metadata, self._active_request_message_id) + reply = _RemoteAgentReply( + stage=stage or self._fallback_stage(), + text=feedback_text, + done=False, + result=None, + event_type="feedback", + metadata=self._feedback_metadata(payload, metadata), + ) + self._push_queue.put_nowait(reply) + self._conta_feedback_event.set() + + if self._timeline is not None: + self._timeline.emit( + "remote_agent_response", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + sse_event="feedback", + stage=reply.stage, + done=False, + text_len=len(reply.text), + has_result=False, + ) + + def _is_conta_feedback_stream( + self, + *, + method: str, + action: str, + payload: Mapping[str, Any] | None, + ) -> bool: + return ( + self._current_agent_name() == "conta" + and str(method or "").upper() == "POST" + and str(action or "").strip().lower() == "chat" + and payload is not None + ) + + def _begin_conta_feedback_stream(self) -> None: + self._conta_feedback_streams += 1 + self._conta_feedback_event.set() + + def _end_conta_feedback_stream(self) -> None: + self._conta_feedback_streams = max(0, self._conta_feedback_streams - 1) + self._conta_feedback_event.set() + + def _add_pending_speech_interruption(self, payload: Dict[str, Any]) -> None: + if self._pending_interrupt is not None: + interruption = dict(self._pending_interrupt) + interruption_key = str( + interruption.pop("_interruption_field", "speech_interruption") + ) + payload[interruption_key] = interruption + + def _fallback_stage(self) -> str: + if self._last_stage and self._last_stage != "INTRO": + return self._last_stage + return self._default_stage + + async def _next_sse_event(self, line_iterator: Any) -> tuple[str, Any] | None: + event_name = "" + data_lines: List[str] = [] + + while True: + try: + raw_line = await line_iterator.__anext__() + except StopAsyncIteration: + if event_name or data_lines: + break + return None + + line = str(raw_line or "") + if not line: + if event_name or data_lines: + break + continue + if line.startswith(":"): + continue + + field, sep, value = line.partition(":") + if not sep: + field = line + value = "" + elif value.startswith(" "): + value = value[1:] + + if field == "event": + event_name = value.strip().lower() + elif field == "data": + data_lines.append(value) + + data = "\n".join(data_lines).strip() + payload = self._decode_message(data) if data else {} + return event_name, payload + + def _coerce_event_payload(self, event_name: str, payload: Any) -> Dict[str, Any]: + if isinstance(payload, str): + coerced = {"type": event_name or "message", "text": payload} + elif not isinstance(payload, Mapping): + coerced = {"type": event_name or "message"} + else: + coerced = dict(payload) + + if "type" not in coerced and event_name: + coerced["type"] = event_name + + return coerced + + def _payload_to_reply( + self, + *, + event_name: str, + payload: Mapping[str, Any], + chunks: List[str], + ) -> tuple[_RemoteAgentReply | None, str, bool]: + self._update_session_id(payload) + self._update_session_actions(payload) + + message_type = str(payload.get("type") or event_name or "").strip().lower() + prefetch_terminal = message_type in _PREFETCH_TERMINAL_TYPES or event_name in _PREFETCH_TERMINAL_TYPES + + error = payload.get("error") + if event_name == "error" or (error and not prefetch_terminal): + raise RuntimeError(str(error or self._extract_reply_text(payload) or "remote agent SSE error")) + + stage = self._extract_stage(payload) + delta = self._extract_delta(payload) + if message_type == "progress" and not delta: + delta = self._extract_reply_text(payload) + if delta: + chunks.append(delta) + if self._timeline is not None: + self._timeline.emit( + "remote_agent_stream_chunk", + agent=self._current_agent_name(), + chunk_len=len(delta), + ) + + if prefetch_terminal or message_type in _STREAM_TYPES: + return None, message_type, prefetch_terminal + + result: Any = None + if "result" in payload: + result = payload.get("result") + elif message_type in {"result", "proactive_result"}: + result = dict(payload) + metadata = self._reply_metadata(payload, result) + + reply_text = self._extract_reply_text(payload) + if self._is_conta_feedback(message_type) or self._is_conta_result_feedback(result): + self._queue_feedback_reply( + payload=payload, + stage=stage or self._fallback_stage(), + text=reply_text, + metadata=metadata, + ) + return None, message_type, prefetch_terminal + + done = ( + final_stop_from_agent_result(result) is not None + or message_type == "done" + or bool(payload.get("done")) + or stage == "DONE" + ) + terminal = ( + bool(payload.get("final")) + or message_type in _TERMINAL_TYPES + or bool(reply_text) + or result is not None + or done + ) + if not terminal: + return None, message_type, prefetch_terminal + + text = (reply_text or "".join(chunks)).strip() + final_stage = stage or self._fallback_stage() + if done: + final_stage = "DONE" + + reply = _RemoteAgentReply( + stage=final_stage, + text=text, + done=done, + result=result, + event_type=message_type, + metadata=metadata, + ) + + if self._timeline is not None: + self._timeline.emit( + "remote_agent_response", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + stage=reply.stage, + done=reply.done, + text_len=len(reply.text), + has_result=reply.result is not None, + ) + logger.info( + "REMOTE_AGENT_SSE_RESPONSE | backend=%s | agent=%s | stage=%s | done=%s | text=%r", + self._backend_label, + self._current_agent_name(), + reply.stage, + reply.done, + reply.text, + ) + + chunks.clear() + return reply, message_type, prefetch_terminal + + def _oferta_done_result(self, payload: Mapping[str, Any]) -> tuple[Dict[str, Any], bool]: + status = str(payload.get("status") or "").strip().lower() + additional = payload.get("additionalInformations") + if not isinstance(additional, Mapping): + additional = {} + + service_status = str(additional.get("service_status") or "").strip().upper() + if status == "transferred": + return ( + { + "type": "transferred", + "content": "", + "status": str(payload.get("status") or "").strip() or "completed", + "additionalInformations": dict(additional), + }, + False, + ) + + if service_status in _OFERTA_NON_FINAL_SERVICE_STATUSES: + return ( + { + "type": service_status.lower(), + "content": "", + "status": str(payload.get("status") or "").strip() or "completed", + "additionalInformations": dict(additional), + }, + False, + ) + + result_type = _OFERTA_SERVICE_STATUS_RESULT_TYPE.get(service_status) + if not result_type: + return ( + { + "type": service_status.lower() if service_status else "completed", + "content": "", + "status": str(payload.get("status") or "").strip() or "completed", + "additionalInformations": dict(additional), + }, + False, + ) + + return ( + { + "type": result_type, + "content": "", + "status": str(payload.get("status") or "").strip() or "completed", + "additionalInformations": dict(additional), + }, + True, + ) + + def _oferta_payload_to_reply( + self, + *, + event_name: str, + payload: Mapping[str, Any], + ) -> _RemoteAgentReply | None: + normalized_event = str(event_name or payload.get("type") or "").strip().lower() + self._update_session_id(payload) + + if normalized_event == "schedule_message": + scheduled = payload.get("scheduledMessage") + if not isinstance(scheduled, Mapping): + scheduled = {} + text = self._extract_text_from_value(scheduled.get("message")) + if not text: + return None + return _RemoteAgentReply( + stage=self._fallback_stage(), + text=text, + done=False, + result=None, + event_type="schedule_message", + metadata={ + "event": "schedule_message", + "scheduledMessage": dict(scheduled), + "payload": dict(payload), + **(self._reply_metadata(payload) or {}), + }, + ) + + if normalized_event == "message": + text = ( + self._extract_text_from_value(payload.get("response")) + or self._extract_reply_text(payload) + ) + if not text: + return None + metadata: Dict[str, Any] = { + "event": "message", + "payload": dict(payload), + } + speech_metadata = self._reply_metadata(payload) + if speech_metadata: + metadata.update(speech_metadata) + if payload.get("sessionId") not in (None, ""): + metadata["sessionId"] = self._string_or_empty(payload.get("sessionId")) + additional = payload.get("additionalInformations") + if isinstance(additional, Mapping): + metadata["additionalInformations"] = dict(additional) + return _RemoteAgentReply( + stage=self._fallback_stage(), + text=text, + done=False, + result=None, + event_type="message", + metadata=metadata, + ) + + if normalized_event == "done": + additional = payload.get("additionalInformations") + if not isinstance(additional, Mapping): + additional = {} + result, finalizes_conversation = self._oferta_done_result(payload) + metadata = { + "event": "done", + "payload": dict(payload), + "additionalInformations": dict(additional), + } + return _RemoteAgentReply( + stage="DONE" if finalizes_conversation else self._fallback_stage(), + text="", + done=finalizes_conversation, + result=result, + event_type="done", + metadata=metadata, + ) + + error = payload.get("error") + if normalized_event == "error" or error: + raise RuntimeError(str(error or self._extract_reply_text(payload) or "remote oferta SSE error")) + + return None + + async def _consume_oferta_execute_stream( + self, + *, + url: str, + payload: Dict[str, Any], + ) -> None: + self._emit_request(method="POST", url=url, action="execute", params={}, payload=payload) + request_started = time.perf_counter() + saw_event = False + saw_done = False + + async with httpx.AsyncClient( + follow_redirects=True, + timeout=self._client_timeout(), + verify=self._verify_tls(), + ) as client: + async with client.stream( + "POST", + url, + params={}, + headers=self._oferta_headers(), + json=payload, + ) as response: + response.raise_for_status() + line_iterator = response.aiter_lines() + while True: + event = await self._next_sse_event(line_iterator) + if event is None: + break + + saw_event = True + event_name, raw_payload = event + event_payload = self._coerce_event_payload(event_name, raw_payload) + if self._timeline is not None: + self._timeline.emit( + "remote_agent_response_raw", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + sse_event=event_name or event_payload.get("type", ""), + response_payload=dict(event_payload), + ) + reply = self._oferta_payload_to_reply( + event_name=event_name, + payload=event_payload, + ) + if reply is not None: + self._reply_with_message_id(reply, payload.get("messageId")) + if self._timeline is not None: + self._timeline.emit( + "remote_agent_response", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + sse_event=reply.event_type, + stage=reply.stage, + done=reply.done, + text_len=len(reply.text), + has_result=reply.result is not None, + ) + logger.info( + "REMOTE_AGENT_SSE_RESPONSE | backend=%s | agent=%s | event=%s | stage=%s | done=%s | duration_ms=%s | text=%r", + self._backend_label, + self._current_agent_name(), + reply.event_type, + reply.stage, + reply.done, + round((time.perf_counter() - request_started) * 1000), + reply.text, + ) + await self._push_queue.put(reply) + if reply.event_type == "done": + saw_done = True + break + + if not saw_event: + raise RuntimeError("Remote oferta SSE stream closed without a usable response") + if not saw_done: + raise RuntimeError("Remote oferta SSE stream closed before done") + + async def _start_oferta_execute_stream_locked(self, payload: Dict[str, Any]) -> _RemoteAgentReply: + await self._cancel_oferta_stream_locked() + self._clear_push_queue() + url = self._resolve_url() + self._offer_stream_task = asyncio.create_task( + self._consume_oferta_execute_stream(url=url, payload=payload), + name="remote_sse_oferta_stream", + ) + reply = await self._next_pushed_reply() + if reply is None: + raise RuntimeError("Remote oferta SSE stream closed without a usable response") + return reply + + async def _cancel_oferta_stream_locked(self) -> None: + task = self._offer_stream_task + if task is None: + return + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + else: + try: + task.exception() + except asyncio.CancelledError: + pass + self._offer_stream_task = None + + async def _consume_sse_response( + self, + line_iterator: Any, + *, + stop_on_first_result: bool, + stop_on_ready: bool, + stop_on_prefetch_terminal: bool, + ) -> list[_RemoteAgentReply]: + chunks: List[str] = [] + replies: list[_RemoteAgentReply] = [] + saw_event = False + + while True: + event = await self._next_sse_event(line_iterator) + if event is None: + break + + saw_event = True + event_name, raw_payload = event + payload = self._coerce_event_payload(event_name, raw_payload) + reply, message_type, prefetch_terminal = self._payload_to_reply( + event_name=event_name, + payload=payload, + chunks=chunks, + ) + if reply is not None: + replies.append(reply) + if stop_on_ready and message_type == "ready": + break + if stop_on_first_result and message_type in {"result", "proactive_result", "done", "final"}: + break + + if stop_on_prefetch_terminal and prefetch_terminal: + break + + if chunks: + replies.append( + _RemoteAgentReply( + stage=self._fallback_stage(), + text="".join(chunks).strip(), + done=False, + result=None, + event_type="progress", + ) + ) + + if not saw_event: + raise RuntimeError("Remote agent SSE stream closed without a usable response") + + return replies + + async def _stream_request( + self, + *, + method: str, + url: str, + params: Mapping[str, Any], + action: str, + payload: Dict[str, Any] | None = None, + stop_on_first_result: bool = False, + stop_on_ready: bool = False, + stop_on_prefetch_terminal: bool = False, + ) -> list[_RemoteAgentReply]: + self._emit_request(method=method, url=url, action=action, params=params, payload=payload) + request_started = time.perf_counter() + feedback_stream = self._is_conta_feedback_stream( + method=method, + action=action, + payload=payload, + ) + request_message_id = self._message_id_from_mapping(params) + previous_request_message_id = self._active_request_message_id + self._active_request_message_id = request_message_id + if feedback_stream: + self._begin_conta_feedback_stream() + try: + async with httpx.AsyncClient( + follow_redirects=True, + timeout=self._client_timeout(), + verify=self._verify_tls(), + ) as client: + kwargs: Dict[str, Any] = { + "params": dict(params), + "headers": self._stream_headers(), + } + if payload is not None: + kwargs["json"] = payload + + stream_cm = client.stream(method, url, **kwargs) + async with stream_cm as response: + response.raise_for_status() + replies = await self._consume_sse_response( + response.aiter_lines(), + stop_on_first_result=stop_on_first_result, + stop_on_ready=stop_on_ready, + stop_on_prefetch_terminal=stop_on_prefetch_terminal, + ) + logger.info( + "REMOTE_AGENT_SSE_REQUEST_DONE | backend=%s | agent=%s | method=%s | action=%s | duration_ms=%s | replies=%s", + self._backend_label, + self._current_agent_name(), + str(method or "").upper(), + action, + round((time.perf_counter() - request_started) * 1000), + len(replies), + ) + if request_message_id: + for reply in replies: + self._reply_with_message_id(reply, request_message_id) + return replies + finally: + self._active_request_message_id = previous_request_message_id + if feedback_stream: + self._end_conta_feedback_stream() + + @staticmethod + def _select_ready_reply(replies: list[_RemoteAgentReply]) -> _RemoteAgentReply | None: + for reply in replies: + if reply.event_type == "ready": + return reply + return replies[0] if replies else None + + @staticmethod + def _select_action_reply(replies: list[_RemoteAgentReply]) -> _RemoteAgentReply: + for reply in replies: + if reply.event_type in {"result", "proactive_result", "done", "final"}: + return reply + for reply in replies: + if reply.event_type != "ready": + return reply + if replies: + return replies[-1] + raise RuntimeError("Remote agent SSE stream closed without a usable response") + + async def prepare( + self, + elegibility: bool, + protocol: str, + ) -> None: + if self._timeline is not None: + self._timeline.emit( + "backend_prepare_started", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + ) + self._elegibility = bool(elegibility) + self._protocol = str(protocol or "") + self._remote_session_id = "" + self._turn_seq = 0 + self._last_stage = "INTRO" + self._pending_interrupt = None + self._pending_events.clear() + self._pending_ready_reply = None + self._session_actions = None + self._conta_feedback_streams = 0 + self._conta_feedback_event.clear() + async with self._session_lock: + await self._cancel_oferta_stream_locked() + self._clear_push_queue() + + if self._is_oferta_agent(): + if self._timeline is not None: + self._timeline.emit( + "backend_prepare_completed", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + remote_session_id="", + ) + return + + async with self._session_lock: + url = self._resolve_url() + replies = await self._stream_request( + method="GET", + url=url, + params=self._stream_open_params(), + action="connect", + stop_on_first_result=False, + stop_on_ready=True, + stop_on_prefetch_terminal=True, + ) + self._pending_ready_reply = self._select_ready_reply(replies) + for reply in replies: + if reply is not self._pending_ready_reply and reply.text: + await self._push_queue.put(reply) + + if self._timeline is not None: + self._timeline.emit( + "backend_prepare_completed", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + remote_session_id=self._remote_session_id, + ) + + async def run(self, user_input: Any) -> BackendReply: + user_text = self._extract_user_text(user_input) + sent_payload = False + + async with self._session_lock: + if self._is_oferta_agent(): + payload = self._build_oferta_execute_payload(user_input) + reply = await self._start_oferta_execute_stream_locked(payload) + sent_payload = True + elif self._pending_ready_reply is not None and not user_text: + reply = self._pending_ready_reply + self._pending_ready_reply = None + elif not user_text: + self._pending_ready_reply = None + reply = _RemoteAgentReply( + stage=self._fallback_stage(), + text="", + done=False, + result=None, + event_type="noop", + ) + else: + self._pending_ready_reply = None + message_id = ( + self._query_message_id(user_input) + if self._current_agent_name() == "conta" + else "" + ) + payload = self._build_turn_payload(user_input, message_id=message_id) + replies = await self._stream_request( + method="POST", + url=self._resolve_url(), + params=self._action_params(user_input, message_id=message_id), + action=str(payload.get("action") or "chat"), + payload=payload, + stop_on_first_result=True, + stop_on_prefetch_terminal=False, + ) + reply = self._select_action_reply(replies) + sent_payload = True + + if sent_payload: + self._clear_pending_state() + self._last_stage = reply.stage or self._last_stage or self._default_stage + return BackendReply( + stage=self._last_stage, + text=reply.text, + done=bool(reply.done) or self._last_stage == "DONE", + export_payload=reply.result, + metadata=reply.metadata, + ) + + async def set_interruption( + self, + interrupted: bool, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + if not interrupted: + self._pending_interrupt = None + if self._timeline is not None: + self._timeline.emit( + "backend_set_interruption", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + interrupted=False, + ) + return + + speech_id = self._string_or_empty(speech_id) + if not speech_id: + self._pending_interrupt = None + if self._timeline is not None: + self._timeline.emit( + "backend_set_interruption", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + interrupted=True, + skipped=True, + listened_text=listened_text, + speech_id="", + reason="missing_speech_id", + ) + return + self._pending_interrupt = { + "speech_id": speech_id, + "heard_text": listened_text or "", + } + if self._timeline is not None: + self._timeline.emit( + "backend_set_interruption", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + interrupted=True, + skipped=bool(skipped), + listened_text=listened_text, + speech_id=speech_id, + ) + + async def set_processing_interruption( + self, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + await self.set_interruption( + True, + listened_text=listened_text, + skipped=skipped, + speech_id=speech_id, + ) + if self._pending_interrupt is not None: + self._pending_interrupt["_interruption_field"] = "processing_interruption" + + async def inject_idle_nudge(self, nudge_text: str) -> None: + text = (nudge_text or "").strip() + if not text: + return + + event = { + "type": "idle_nudge", + "text": text, + } + # Cada frase de inatividade substitui a anterior: o cliente responde ao + # que ouviu por ultimo, e as intermediarias so empilham falas do agente + # no historico remoto -- inclusive o aviso de encerramento, que passa a + # constar como dito logo antes de a conversa seguir normalmente. + if self._pending_events and self._pending_events[-1].get("type") == "idle_nudge": + self._pending_events[-1] = event + else: + self._pending_events.append(event) + if self._timeline is not None: + self._timeline.emit( + "backend_idle_nudge_buffered", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + text=text, + ) + + async def end_service_once(self) -> BackendReply: + async with self._end_lock: + if self._ended: + return self._end_reply or BackendReply(stage="DONE", done=True, export_payload=[]) + + self._ended = True + if self._is_oferta_agent(): + async with self._session_lock: + await self._cancel_oferta_stream_locked() + self._end_reply = BackendReply(stage="DONE", done=True, export_payload=[]) + return self._end_reply + + if self._timeline is not None: + self._timeline.emit( + "backend_end_started", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + ) + try: + async with self._session_lock: + if not self._session_supports_action("end"): + reply = _RemoteAgentReply( + stage="DONE", + text="", + done=True, + result=[], + event_type="done", + ) + else: + message_id = ( + self._query_message_id() + if self._current_agent_name() == "conta" + else "" + ) + payload = self._build_end_payload(message_id=message_id) + try: + replies = await self._stream_request( + method="POST", + url=self._resolve_url(), + params=self._action_params(message_id=message_id), + action="end", + payload=payload, + stop_on_first_result=True, + stop_on_prefetch_terminal=False, + ) + reply = self._select_action_reply(replies) + except RuntimeError as exc: + if not self._is_unsupported_action_error(exc, "end"): + raise + reply = _RemoteAgentReply( + stage="DONE", + text="", + done=True, + result=[], + event_type="done", + ) + except Exception: + logger.exception("[remote-agent] end_service SSE falhou") + if self._timeline is not None: + self._timeline.emit( + "backend_end_failed", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + ) + self._end_reply = BackendReply(stage="DONE", done=True, export_payload=[]) + return self._end_reply + + self._clear_pending_state() + stage = (reply.stage or self._last_stage or self._default_stage).upper() + self._last_stage = stage + self._end_reply = BackendReply( + stage=stage, + text=reply.text, + done=bool(reply.done) or stage == "DONE", + export_payload=reply.result if reply.result is not None else [], + metadata=reply.metadata, + ) + if self._timeline is not None: + self._timeline.emit( + "backend_end_completed", + backend=self._backend_label, + backend_family="remote_sse", + agent=self._current_agent_name(), + result_type=type(self._end_reply.export_payload).__name__, + ) + return self._end_reply diff --git a/src/app/livekit/adapters/remote_agent_ws_adapter.py b/src/app/livekit/adapters/remote_agent_ws_adapter.py new file mode 100644 index 0000000..433541f --- /dev/null +++ b/src/app/livekit/adapters/remote_agent_ws_adapter.py @@ -0,0 +1,1577 @@ +from __future__ import annotations + +import asyncio +import importlib +import json +import logging +import os +import re +import socket +import time +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import Any, Dict, List, Mapping, Optional +from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit + +from app.livekit.adapters.agent_backend import BackendReply +from app.livekit.policies.agent_finalization import final_stop_from_agent_result +from app.utils.logging import setup_minimal_logging +from app.utils.turn_ids import next_turn_message_id + +logger = setup_minimal_logging() + +_STREAM_TYPES = {"chunk", "delta", "partial", "token"} +_TERMINAL_TYPES = {"done", "final", "message", "output", "response", "result"} +_TEXT_KEYS = ("text", "response", "output", "message", "content") +_USER_TEXT_KEYS = ("text", "transcript", "utterance", "message", "content") +_CONTA_USER_RESPONSE_MESSAGE_TYPES = {"ready", "final"} + + +@dataclass(slots=True) +class _RemoteAgentReply: + stage: str + text: str + done: bool + result: Any + metadata: Any = None + + +class RemoteAgentWSAdapter: + def __init__( + self, + *, + intro: str, + url: str, + backend_label: str = "remote_ws", + url_by_agent: Optional[Dict[str, str]] = None, + request_context: Optional[Dict[str, Any]] = None, + timeline: Any | None = None, + streaming: bool = False, + default_stage: str = "PRESENTATION", + open_timeout_s: float = 10.0, + read_timeout_s: float = 45.0, + write_timeout_s: float = 10.0, + close_timeout_s: float = 10.0, + max_message_bytes: int = 1024 * 1024, + ) -> None: + self._intro = intro + self._url = url + self._backend_label = str(backend_label or "remote_ws").strip().lower() + self._url_by_agent = { + self._normalize_agent_name(key): value.strip() + for key, value in (url_by_agent or {}).items() + if value and value.strip() + } + self._request_context = dict(request_context or {}) + self._timeline = timeline + self._streaming = streaming + self._default_stage = (default_stage or "PRESENTATION").strip().upper() + self._open_timeout_s = open_timeout_s + self._read_timeout_s = read_timeout_s + self._write_timeout_s = write_timeout_s + self._close_timeout_s = close_timeout_s + self._max_message_bytes = max_message_bytes + self._instance_name = (os.getenv("HOSTNAME") or socket.gethostname() or "").strip() or "-" + + self._elegibility = True + self._protocol = "" + self._last_stage = "INTRO" + self._pending_interrupt: Optional[Dict[str, Any]] = None + self._pending_events: List[Dict[str, Any]] = [] + + self._end_lock = asyncio.Lock() + self._ended = False + self._end_reply: Optional[BackendReply] = None + self._session_lock = asyncio.Lock() + self._session_cm: Any | None = None + self._session_websocket: Any | None = None + self._pending_ready_reply: Optional[_RemoteAgentReply] = None + self._reader_task: asyncio.Task[Any] | None = None + self._response_waiter: asyncio.Future[_RemoteAgentReply] | None = None + self._push_queue: asyncio.Queue[_RemoteAgentReply] = asyncio.Queue() + self._reader_error: BaseException | None = None + self._terminal_reply: Optional[_RemoteAgentReply] = None + self._session_actions: set[str] | None = None + self._query_message_seq = 0 + self._connection_message_id = "" + self._active_request_message_id = "" + + @staticmethod + def _normalize_agent_name(value: Any) -> str: + raw = str(value or "").strip().lower() + aliases = { + "conta": "conta", + "contas": "conta", + "ofert": "oferta", + "oferta": "oferta", + "ofertas": "oferta", + "cobra": "cobranca", + "cobranca": "cobranca", + "cobrança": "cobranca", + "cobrancas": "cobranca", + "cobranças": "cobranca", + } + return aliases.get(raw, raw) + + @staticmethod + def _string_or_empty(value: Any) -> str: + return str(value or "").strip() + + @classmethod + def _message_id_from_mapping(cls, payload: Mapping[str, Any] | None) -> str: + if not isinstance(payload, Mapping): + return "" + + for key in ("message_id", "messageId"): + value = payload.get(key) + if value not in (None, ""): + return cls._string_or_empty(value) + + for container_key in ("metadata", "data", "payload", "context"): + metadata = payload.get(container_key) + if not isinstance(metadata, Mapping): + continue + for key in ("message_id", "messageId"): + value = metadata.get(key) + if value not in (None, ""): + return cls._string_or_empty(value) + + return "" + + @classmethod + def _metadata_with_message_id(cls, metadata: Any, message_id: Any) -> Dict[str, Any] | None: + normalized_message_id = cls._string_or_empty(message_id) + if not normalized_message_id: + return dict(metadata) if isinstance(metadata, Mapping) else metadata + + merged: Dict[str, Any] = dict(metadata) if isinstance(metadata, Mapping) else {} + merged.setdefault("message_id", normalized_message_id) + return merged + + @classmethod + def _reply_with_message_id( + cls, + reply: _RemoteAgentReply, + message_id: Any, + ) -> _RemoteAgentReply: + metadata = cls._metadata_with_message_id(reply.metadata, message_id) + if metadata is not reply.metadata: + reply.metadata = metadata + return reply + + def _import_websockets(self): + module = importlib.import_module("websockets") + connect = getattr(module, "connect", None) + if callable(connect): + return connect + + client_module = getattr(module, "client", None) + connect = getattr(client_module, "connect", None) + if callable(connect): + return connect + + raise RuntimeError("websockets.connect is not available") + + def _serialize_message(self, payload: Dict[str, Any]) -> str: + return json.dumps(payload, ensure_ascii=False) + + def _decode_message(self, raw_message: Any) -> Any: + if isinstance(raw_message, bytes): + raw_message = raw_message.decode("utf-8") + if not isinstance(raw_message, str): + raise RuntimeError(f"Unsupported websocket payload type: {type(raw_message)!r}") + + message = raw_message.strip() + if not message: + return {} + + try: + return json.loads(message) + except json.JSONDecodeError: + return message + + def _extract_text_from_value(self, value: Any) -> str: + if isinstance(value, str): + return value.strip() + if isinstance(value, Mapping): + for key in ("result", * _TEXT_KEYS): + text = self._extract_text_from_value(value.get(key)) + if text: + return text + return "" + if isinstance(value, list): + for item in value: + text = self._extract_text_from_value(item) + if text: + return text + return "" + + def _extract_user_text(self, user_input: Any) -> str: + if isinstance(user_input, str): + return user_input.strip() + + if isinstance(user_input, Mapping): + for key in _USER_TEXT_KEYS: + text = self._extract_text_from_value(user_input.get(key)) + if text: + return text + + for value in user_input.values(): + text = self._extract_text_from_value(value) + if text: + return text + + if isinstance(user_input, list): + for item in user_input: + text = self._extract_user_text(item) + if text: + return text + + return str(user_input).strip() + + def _copy_context(self) -> Dict[str, Any]: + return { + "intro": self._intro, + "elegibility": self._elegibility, + "streaming": self._streaming, + } + + def _current_timestamp(self) -> str: + return datetime.now(timezone.utc).isoformat() + + def _request_field(self, *keys: str) -> str: + for key in keys: + value = self._request_context.get(key) + if value not in (None, ""): + return self._string_or_empty(value) + + lowered = {str(key).lower(): value for key, value in self._request_context.items()} + for key in keys: + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return self._string_or_empty(value) + + return "" + + def _current_agent_name(self) -> str: + return self._normalize_agent_name( + self._request_field("agent", "Agent", "agente") + ) + + def _resolve_url(self) -> str: + agent_name = self._current_agent_name() + if agent_name: + routed_url = self._url_by_agent.get(agent_name, "") + if routed_url: + return routed_url + return self._url + + def _current_invoice_number(self) -> str: + return self._request_field( + "current_invoice_number", + "currentInvoiceNumber", + "ID_FATURA", + "idFatura", + "id_fatura", + "IdFatura", + ) + + def _current_msisdn(self) -> str: + return self._request_field("msisdn", "GSM", "gsm", "NUM_TELEFONE") + + def _current_ani(self) -> str: + return self._request_field("ANI", "ani") + + def _current_protocol_number(self) -> str: + return self._request_field("protocol_id", "protocolId", "protocolo", "protocolNumber", "protocol") + + def _current_session_id(self) -> str: + return self._request_field("session_id", "sessionId") + + def _next_query_message_id(self) -> str: + source = dict(self._request_context) + if self._protocol: + source.setdefault("protocol", self._protocol) + return next_turn_message_id(source) + + def _current_connection_message_id(self) -> str: + explicit_message_id = self._request_field("message_id", "messageId") + if explicit_message_id: + return explicit_message_id + + if not self._connection_message_id: + self._connection_message_id = self._next_query_message_id() + return self._connection_message_id + + def _current_channel_id(self) -> str: + channel_id = self._request_field("channelId", "channel_id") + if channel_id: + return channel_id + return "ura" + + def _current_ura_call_id(self) -> str: + return self._request_field("uraCallId", "ura_call_id", "callIdGed") + + def _current_channel(self) -> str: + channel = self._request_field("channel", "Channel", "canal") + if channel: + return channel + + configured = self._string_or_empty( + os.getenv("REMOTE_AGENT_CHANNEL") + or os.getenv("REMOTE_AGENT_DEFAULT_CHANNEL") + ) + return configured or "SUPERVISOR" + + def _connection_query_params(self) -> Dict[str, str]: + if self._current_agent_name() != "conta": + return {} + + params: Dict[str, str] = {} + msisdn = self._current_msisdn() + if msisdn: + params["msisdn"] = msisdn + + current_invoice_number = self._current_invoice_number() + if current_invoice_number: + params["current_invoice_number"] = current_invoice_number + + ani = self._current_ani() + if ani: + params["ani"] = ani + + protocol_id = self._current_protocol_number() + if protocol_id: + params["protocol_id"] = protocol_id + + session_id = self._current_session_id() + if session_id: + params["session_id"] = session_id + + message_id = self._current_connection_message_id() + if message_id: + params["message_id"] = message_id + + channel_id = self._current_channel_id() + if channel_id: + params["channelId"] = channel_id + + ura_call_id = self._current_ura_call_id() + if ura_call_id: + params["uraCallId"] = ura_call_id + + return params + + @staticmethod + def _append_query_params(url: str, params: Mapping[str, Any]) -> str: + filtered_params = { + str(key): str(value).strip() + for key, value in params.items() + if value not in (None, "") and str(value).strip() + } + if not filtered_params: + return url + + parsed = urlsplit(url) + query_params = dict(parse_qsl(parsed.query, keep_blank_values=True)) + query_params.update(filtered_params) + return urlunsplit( + ( + parsed.scheme, + parsed.netloc, + parsed.path, + urlencode(query_params), + parsed.fragment, + ) + ) + + def _uses_persistent_session(self) -> bool: + return self._current_agent_name() == "conta" + + def _connection_url(self) -> str: + return self._append_query_params( + self._resolve_url(), + self._connection_query_params(), + ) + + def _connection_kwargs(self) -> Dict[str, Any]: + return { + "open_timeout": self._open_timeout_s, + "close_timeout": self._close_timeout_s, + "max_size": self._max_message_bytes, + } + + @staticmethod + def _compact_log_value(value: Any) -> str: + text = str(value or "").strip() + if not text: + return "-" + if len(text) > 240: + return f"{text[:237]}..." + return text.replace("\n", "\\n") + + def _log_ws_event(self, event: str, *, level: int = logging.INFO, **fields: Any) -> None: + base_fields: Dict[str, Any] = { + "backend": self._backend_label, + "agent": self._current_agent_name() or "-", + "persistent": int(self._uses_persistent_session()), + "instance": self._instance_name, + } + base_fields.update(fields) + + ordered_keys = ( + "backend", + "agent", + "persistent", + "instance", + "url", + "host", + "action", + "stage", + "protocol", + "duration_ms", + "text_len", + "text_preview", + "payload_bytes", + "done", + "result_type", + "error_type", + "error", + ) + parts = [] + used_keys: set[str] = set() + for key in ordered_keys: + if key not in base_fields: + continue + value = base_fields[key] + if value in (None, ""): + continue + parts.append(f"{key}={self._compact_log_value(value)}") + used_keys.add(key) + + for key, value in base_fields.items(): + if key in used_keys or value in (None, ""): + continue + parts.append(f"{key}={self._compact_log_value(value)}") + + logger.log(level, "REMOTE_AGENT_WS_%s | %s", event, " | ".join(parts)) + + def _log_ws_failure(self, event: str, error: BaseException, **fields: Any) -> None: + self._log_ws_event( + event, + level=logging.ERROR, + error_type=type(error).__name__, + error=str(error), + **fields, + ) + + def _payload_metadata(self, payload: Mapping[str, Any]) -> Dict[str, Any]: + body = payload.get("payload") + if not isinstance(body, Mapping): + body = payload + + action = str(payload.get("action") or body.get("action") or "").strip().lower() + stage = str(body.get("stage") or payload.get("stage") or self._fallback_stage()).strip().upper() + protocol = str(body.get("protocol") or payload.get("protocol") or self._protocol or "").strip() + if "message" in body: + user_input = body.get("message") or "" + elif "text" in body: + user_input = body.get("text") or "" + else: + user_input = "" + user_text = self._extract_user_text(user_input) + return { + "action": action or "chat", + "stage": stage or self._fallback_stage(), + "protocol": protocol or "-", + "text_len": len(user_text), + } + + @staticmethod + def _url_host(url: str) -> str: + parsed = urlsplit(str(url or "")) + return (parsed.hostname or "").strip() + + def _update_session_actions(self, payload: Mapping[str, Any]) -> None: + if "actions" not in payload: + return + + raw_actions = payload.get("actions") + parsed_actions: set[str] = set() + if isinstance(raw_actions, list): + for item in raw_actions: + action = str(item or "").strip().lower() + if action: + parsed_actions.add(action) + elif isinstance(raw_actions, str): + action = raw_actions.strip().lower() + if action: + parsed_actions.add(action) + + self._session_actions = parsed_actions + + def _persistent_session_supports_action(self, action: str) -> bool: + normalized_action = str(action or "").strip().lower() + if not normalized_action: + return False + + if not self._uses_persistent_session(): + return True + + if self._session_actions is None: + return normalized_action != "end" + + return normalized_action in self._session_actions + + @staticmethod + def _is_unsupported_action_error(error: BaseException, action: str) -> bool: + action_name = str(action or "").strip().lower() + if not action_name: + return False + message = str(error or "").strip().lower() + return "não suportada" in message and action_name in message + + def _emit_request(self, *, url: str, payload: Dict[str, Any]) -> None: + if self._timeline is not None: + self._timeline.emit( + "remote_agent_request", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + url=url, + action=payload.get("action", ""), + payload=payload, + ) + try: + payload_bytes = len(self._serialize_message(payload).encode("utf-8")) + except Exception: + payload_bytes = None + self._log_ws_event( + "REQUEST", + url=url, + host=self._url_host(url), + payload_bytes=payload_bytes, + **self._payload_metadata(payload), + ) + + async def _open_persistent_session_locked(self, *, capture_ready: bool) -> None: + if self._session_websocket is not None: + return + + connect = self._import_websockets() + url = self._connection_url() + self._log_ws_event( + "CONNECT_OPEN", + url=url, + host=self._url_host(url), + protocol=self._protocol or "-", + ) + session_cm = connect(url, **self._connection_kwargs()) + try: + websocket = await session_cm.__aenter__() + except Exception as exc: + self._log_ws_failure( + "CONNECT_FAIL", + exc, + url=url, + host=self._url_host(url), + protocol=self._protocol or "-", + ) + raise + self._session_cm = session_cm + self._session_websocket = websocket + self._reader_error = None + self._terminal_reply = None + self._clear_push_queue() + self._log_ws_event( + "CONNECT_OK", + url=url, + host=self._url_host(url), + protocol=self._protocol or "-", + ) + + ready_waiter: asyncio.Future[_RemoteAgentReply] | None = None + if capture_ready: + ready_waiter = self._create_response_waiter_locked() + + self._start_reader_task_locked() + + if ready_waiter is None: + return + + try: + self._pending_ready_reply = await asyncio.wait_for( + ready_waiter, + timeout=self._read_timeout_s, + ) + except asyncio.TimeoutError: + self._log_ws_event( + "READY_TIMEOUT", + url=url, + host=self._url_host(url), + protocol=self._protocol or "-", + ) + if self._response_waiter is ready_waiter: + self._response_waiter = None + if ready_waiter is not None and not ready_waiter.done(): + ready_waiter.cancel() + await self._close_persistent_session_locked() + raise + except Exception as exc: + self._log_ws_failure( + "READY_FAIL", + exc, + url=url, + host=self._url_host(url), + protocol=self._protocol or "-", + ) + await self._close_persistent_session_locked() + raise + + async def _close_persistent_session_locked(self) -> None: + reader_task = self._reader_task + self._reader_task = None + waiter = self._response_waiter + self._response_waiter = None + session_cm = self._session_cm + self._session_cm = None + self._session_websocket = None + self._pending_ready_reply = None + self._reader_error = None + self._terminal_reply = None + self._session_actions = None + self._query_message_seq = 0 + self._connection_message_id = "" + self._active_request_message_id = "" + self._clear_push_queue() + + if waiter is not None and not waiter.done(): + waiter.cancel() + + if reader_task is not None and not reader_task.done(): + reader_task.cancel() + await asyncio.gather(reader_task, return_exceptions=True) + + if session_cm is None: + return + + await session_cm.__aexit__(None, None, None) + + async def _send_persistent_payload_locked(self, payload: Dict[str, Any]) -> _RemoteAgentReply: + websocket = self._session_websocket + if websocket is None: + raise RuntimeError("Persistent websocket session is not connected") + + waiter = self._create_response_waiter_locked() + url = self._connection_url() + self._emit_request(url=url, payload=payload) + request_started = time.perf_counter() + request_message_id = self._message_id_from_mapping(payload) + previous_request_message_id = self._active_request_message_id + self._active_request_message_id = request_message_id + try: + await asyncio.wait_for( + websocket.send(self._serialize_message(payload)), + timeout=self._write_timeout_s, + ) + except Exception as exc: + self._log_ws_failure( + "SEND_FAIL", + exc, + url=url, + host=self._url_host(url), + **self._payload_metadata(payload), + ) + if self._response_waiter is waiter: + self._response_waiter = None + if not waiter.done(): + waiter.cancel() + self._active_request_message_id = previous_request_message_id + raise + try: + reply = await asyncio.wait_for(waiter, timeout=self._read_timeout_s) + if request_message_id: + self._reply_with_message_id(reply, request_message_id) + done_fields = { + **self._payload_metadata(payload), + "duration_ms": round((time.perf_counter() - request_started) * 1000), + "stage": reply.stage, + "text_len": len(reply.text), + "text_preview": reply.text, + "done": int(reply.done), + "result_type": type(reply.result).__name__ if reply.result is not None else "-", + } + self._log_ws_event( + "TURN_DONE", + url=url, + host=self._url_host(url), + **done_fields, + ) + return reply + except asyncio.TimeoutError: + self._log_ws_event( + "RESPONSE_TIMEOUT", + url=url, + host=self._url_host(url), + duration_ms=round((time.perf_counter() - request_started) * 1000), + **self._payload_metadata(payload), + ) + if self._response_waiter is waiter: + self._response_waiter = None + if not waiter.done(): + waiter.cancel() + raise + finally: + self._active_request_message_id = previous_request_message_id + + async def _run_with_persistent_session(self, user_input: Any) -> tuple[_RemoteAgentReply, bool]: + user_text = self._extract_user_text(user_input) + + async with self._session_lock: + await self._open_persistent_session_locked(capture_ready=self._session_websocket is None) + + if self._pending_ready_reply is not None: + if not user_text: + reply = self._pending_ready_reply + self._pending_ready_reply = None + return reply, False + self._pending_ready_reply = None + + try: + reply = await self._send_persistent_payload_locked(self._build_turn_payload(user_input)) + except Exception: + await self._close_persistent_session_locked() + raise + + return reply, True + + def _clear_push_queue(self) -> None: + while True: + try: + self._push_queue.get_nowait() + except asyncio.QueueEmpty: + return + + def _create_response_waiter_locked(self) -> asyncio.Future[_RemoteAgentReply]: + current = self._response_waiter + if current is not None and not current.done(): + raise RuntimeError("Another persistent websocket response is already pending") + + waiter: asyncio.Future[_RemoteAgentReply] = asyncio.get_running_loop().create_future() + self._response_waiter = waiter + return waiter + + def _start_reader_task_locked(self) -> None: + if self._reader_task is not None and not self._reader_task.done(): + return + websocket = self._session_websocket + if websocket is None: + raise RuntimeError("Persistent websocket session is not connected") + self._reader_task = asyncio.create_task( + self._reader_loop(websocket), + name="remote_agent_ws_reader", + ) + + async def _reader_loop(self, websocket: Any) -> None: + try: + while True: + reply = await self._recv_reply(websocket, timeout_s=None) + waiter = self._response_waiter + if waiter is not None and not waiter.done(): + self._response_waiter = None + waiter.set_result(reply) + else: + await self._push_queue.put(reply) + + if reply.done: + self._terminal_reply = reply + return + except asyncio.CancelledError: + raise + except Exception as exc: + self._reader_error = exc + self._log_ws_failure( + "READER_FAIL", + exc, + url=self._connection_url(), + host=self._url_host(self._connection_url()), + protocol=self._protocol or "-", + ) + waiter = self._response_waiter + if waiter is not None and not waiter.done(): + self._response_waiter = None + waiter.set_exception(exc) + + def supports_server_push(self) -> bool: + return self._uses_persistent_session() + + def supports_inflight_backend_push(self) -> bool: + return self._current_agent_name() == "conta" + + async def wait_for_server_push(self) -> BackendReply | None: + if not self._uses_persistent_session(): + return None + + while True: + try: + reply = self._push_queue.get_nowait() + except asyncio.QueueEmpty: + reply = None + + if reply is not None: + self._last_stage = reply.stage or self._last_stage or self._default_stage + return BackendReply( + stage=self._last_stage, + text=reply.text, + done=bool(reply.done) or self._last_stage == "DONE", + export_payload=reply.result, + metadata=reply.metadata, + ) + + if self._reader_error is not None: + raise RuntimeError("Persistent websocket reader failed") from self._reader_error + + reader_task = self._reader_task + if reader_task is None or reader_task.done(): + return None + + get_task = asyncio.create_task(self._push_queue.get(), name="remote_agent_ws_push_wait") + try: + done, pending = await asyncio.wait( + {get_task, reader_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + except asyncio.CancelledError: + get_task.cancel() + await asyncio.gather(get_task, return_exceptions=True) + raise + + if get_task in done: + reply = get_task.result() + self._last_stage = reply.stage or self._last_stage or self._default_stage + return BackendReply( + stage=self._last_stage, + text=reply.text, + done=bool(reply.done) or self._last_stage == "DONE", + export_payload=reply.result, + metadata=reply.metadata, + ) + + get_task.cancel() + await asyncio.gather(get_task, return_exceptions=True) + + if self._reader_error is not None: + raise RuntimeError("Persistent websocket reader failed") from self._reader_error + + if reader_task.done() and self._push_queue.empty(): + return None + + def _base_payload(self) -> Dict[str, Any]: + payload = { + "timestamp": self._current_timestamp(), + "agent": self._current_agent_name(), + "RouterCallKeyDay": self._request_field("RouterCallKeyDay", "router_call_key_day", "routerCallKeyDay"), + "RouterCallKey": self._request_field("RouterCallKey", "router_call_key", "routerCallKey"), + "ANI": self._request_field("ANI", "ani"), + "GSM": self._request_field("GSM", "gsm", "NUM_TELEFONE"), + "callIdGed": self._request_field("callIdGed"), + } + + if payload["agent"] == "conta": + id_fatura = self._request_field("ID_FATURA", "id_fatura", "IdFatura") + if id_fatura: + payload["ID_FATURA"] = id_fatura + + return payload + + def _build_conta_turn_payload(self, user_input: Any) -> Dict[str, Any]: + payload = { + "message": self._extract_user_text(user_input), + "channel": self._current_channel(), + "msisdn": self._current_msisdn(), + } + message_id = self._message_id_from_mapping(user_input) + if message_id: + payload["message_id"] = message_id + + current_invoice_number = self._current_invoice_number() + if current_invoice_number: + payload["current_invoice_number"] = current_invoice_number + + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + + return { + "action": "chat", + "payload": payload, + } + + def _build_turn_payload(self, user_input: Any) -> Dict[str, Any]: + if self._current_agent_name() == "conta": + return self._build_conta_turn_payload(user_input) + + payload = self._base_payload() + payload["text"] = self._extract_user_text(user_input) + payload["protocol"] = self._protocol + payload["stage"] = self._last_stage + message_id = self._message_id_from_mapping(user_input) + if message_id: + payload["message_id"] = message_id + + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + + return payload + + def _build_end_payload(self) -> Dict[str, Any]: + if self._current_agent_name() == "conta": + payload = self._base_payload() + payload["protocol"] = self._protocol + payload["stage"] = self._last_stage + payload["msisdn"] = self._current_msisdn() + payload["channel"] = self._current_channel() + + current_invoice_number = self._current_invoice_number() + if current_invoice_number: + payload["current_invoice_number"] = current_invoice_number + + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + + return { + "action": "end", + "payload": payload, + } + + payload = self._base_payload() + payload["type"] = "end" + payload["protocol"] = self._protocol + payload["stage"] = self._last_stage + + self._add_pending_speech_interruption(payload) + if self._pending_events: + payload["events"] = list(self._pending_events) + + return payload + + def _clear_pending_state(self) -> None: + self._pending_interrupt = None + self._pending_events.clear() + + def _extract_stage(self, payload: Mapping[str, Any]) -> str: + stage = payload.get("stage") + if isinstance(stage, str) and stage.strip(): + return stage.strip().upper() + return "" + + def _extract_delta(self, payload: Mapping[str, Any]) -> str: + for key in ("delta", "chunk", "token"): + value = payload.get(key) + if isinstance(value, str) and value.strip(): + return value + return "" + + def _extract_reply_text(self, payload: Mapping[str, Any]) -> str: + message_type = str(payload.get("type") or "").strip().lower() + if message_type == "final": + text = self._extract_text_from_value(payload.get("content")) + if text: + return text + + for key in ("result", *_TEXT_KEYS): + value = payload.get(key) + text = self._extract_text_from_value(value) + if text: + return text + return "" + + @staticmethod + def _bool_or_none(value: Any) -> bool | None: + if isinstance(value, bool): + return value + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"1", "true", "yes", "on"}: + return True + if normalized in {"0", "false", "no", "off"}: + return False + return None + + def _reply_metadata(self, payload: Mapping[str, Any], result: Any = None) -> Dict[str, Any] | None: + raw_metadata = payload.get("metadata") + metadata: Dict[str, Any] = dict(raw_metadata) if isinstance(raw_metadata, Mapping) else {} + + sources: list[Mapping[str, Any]] = [] + if isinstance(result, Mapping): + sources.append(result) + sources.append(payload) + + for source in sources: + message_id = self._message_id_from_mapping(source) + if message_id: + metadata["message_id"] = message_id + break + + for source in sources: + speech_id = self._string_or_empty(source.get("speech_id")) + if speech_id: + metadata["speech_id"] = speech_id + break + + for source in sources: + interruptible = self._bool_or_none(source.get("is_interruptible")) + if interruptible is not None: + metadata["is_interruptible"] = interruptible + break + + if self._current_agent_name() == "conta": + message_type = str(payload.get("type") or "").strip().lower() + result_type = "" + if isinstance(result, Mapping): + result_type = str(result.get("type") or "").strip().lower() + + expects_user_response = ( + result_type == "final" + or (not result_type and message_type in _CONTA_USER_RESPONSE_MESSAGE_TYPES) + ) + terminal_result = final_stop_from_agent_result(result) is not None + protected_speech = ( + message_type in {"ready", "feedback"} + or terminal_result + or (bool(result_type) and result_type != "final") + ) + + if message_type: + metadata["agent_message_type"] = message_type + if result_type: + metadata["agent_result_type"] = result_type + metadata["expects_user_response"] = expects_user_response + metadata["drop_user_input_while_speaking"] = protected_speech + if protected_speech: + metadata["is_interruptible"] = False + + return metadata or None + + def _is_conta_feedback(self, message_type: str) -> bool: + return self._current_agent_name() == "conta" and message_type == "feedback" + + def _is_conta_result_feedback(self, result: Any) -> bool: + return ( + self._current_agent_name() == "conta" + and isinstance(result, Mapping) + and str(result.get("type") or "").strip().lower() == "feedback" + ) + + def _feedback_metadata( + self, + payload: Mapping[str, Any], + base_metadata: Any = None, + ) -> Dict[str, Any]: + metadata: Dict[str, Any] = {} + if isinstance(base_metadata, Mapping): + metadata.update(base_metadata) + metadata["event"] = "feedback" + metadata["payload"] = dict(payload) + return metadata + + def _queue_feedback_reply( + self, + *, + payload: Mapping[str, Any], + stage: str, + text: str, + metadata: Any = None, + ) -> None: + feedback_text = (text or "").strip() + if not feedback_text: + return + + if self._current_agent_name() == "conta": + metadata = self._metadata_with_message_id( + metadata, + self._active_request_message_id + or self._connection_message_id + or self._request_field("message_id", "messageId"), + ) + reply = _RemoteAgentReply( + stage=stage or self._fallback_stage(), + text=feedback_text, + done=False, + result=None, + metadata=self._feedback_metadata(payload, metadata), + ) + self._push_queue.put_nowait(reply) + + if self._timeline is not None: + self._timeline.emit( + "remote_agent_response", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + message_type="feedback", + stage=reply.stage, + done=False, + text_len=len(reply.text), + has_result=False, + ) + + def _add_pending_speech_interruption(self, payload: Dict[str, Any]) -> None: + if self._pending_interrupt is not None: + interruption = dict(self._pending_interrupt) + interruption_key = str( + interruption.pop("_interruption_field", "speech_interruption") + ) + payload[interruption_key] = interruption + + def _fallback_stage(self) -> str: + if self._last_stage and self._last_stage != "INTRO": + return self._last_stage + return self._default_stage + + async def _recv_reply(self, websocket: Any, timeout_s: float | None = None) -> _RemoteAgentReply: + chunks: List[str] = [] + final_text = "" + final_stage = "" + done = False + result: Any = None + metadata: Any = None + saw_message = False + + while True: + try: + recv_coro = websocket.recv() + if timeout_s is None: + raw_message = await recv_coro + else: + raw_message = await asyncio.wait_for(recv_coro, timeout=timeout_s) + except asyncio.TimeoutError: + if saw_message or chunks or final_text or result is not None: + break + raise + except asyncio.CancelledError: + raise + except Exception: + if saw_message or chunks or final_text or result is not None: + break + raise + + saw_message = True + message = self._decode_message(raw_message) + + if isinstance(message, str): + final_text = message + break + + if not isinstance(message, Mapping): + continue + + self._update_session_actions(message) + + error = message.get("error") + if error: + raise RuntimeError(str(error)) + + message_type = str(message.get("type") or "").strip().lower() + stage = self._extract_stage(message) + if stage and not self._is_conta_feedback(message_type): + final_stage = stage + + previous_result = result + if "result" in message: + result = message.get("result") + message_metadata = self._reply_metadata(message, result) + + delta = self._extract_delta(message) + if delta: + chunks.append(delta) + if self._timeline is not None: + self._timeline.emit( + "remote_agent_stream_chunk", + agent=self._current_agent_name(), + chunk_len=len(delta), + ) + + reply_text = self._extract_reply_text(message) + + if self._is_conta_feedback(message_type) or self._is_conta_result_feedback(result): + self._queue_feedback_reply( + payload=message, + stage=stage or final_stage or self._fallback_stage(), + text=reply_text, + metadata=message_metadata, + ) + result = previous_result + continue + + if message_metadata is not None: + metadata = message_metadata + + if reply_text: + final_text = reply_text + + if final_stop_from_agent_result(result) is not None: + done = True + if message_type == "done": + done = True + if done or bool(message.get("done")) or final_stage == "DONE": + done = True + break + + if message_type in _STREAM_TYPES: + continue + + if bool(message.get("final")) or message_type in _TERMINAL_TYPES or reply_text: + break + + text = (final_text or "".join(chunks)).strip() + stage = final_stage or self._fallback_stage() + if done: + stage = "DONE" + + if not text and result is None: + raise RuntimeError("Remote agent websocket closed without a usable response") + + if self._timeline is not None: + self._timeline.emit( + "remote_agent_response", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + stage=stage, + done=done, + text_len=len(text), + has_result=result is not None, + ) + self._log_ws_event( + "RESPONSE", + url=self._connection_url(), + host=self._url_host(self._connection_url()), + stage=stage, + protocol=self._protocol or "-", + text_len=len(text), + text_preview=text, + done=int(done), + result_type=type(result).__name__ if result is not None else "-", + ) + + reply = _RemoteAgentReply( + stage=stage, + text=text, + done=done, + result=result, + metadata=metadata, + ) + if self._current_agent_name() == "conta": + self._reply_with_message_id( + reply, + self._active_request_message_id + or self._connection_message_id + or self._request_field("message_id", "messageId"), + ) + return reply + + async def _request(self, payload: Dict[str, Any]) -> _RemoteAgentReply: + connect = self._import_websockets() + url = self._connection_url() + self._log_ws_event( + "CONNECT_OPEN", + url=url, + host=self._url_host(url), + protocol=self._protocol or "-", + ) + request_message_id = self._message_id_from_mapping(payload) + previous_request_message_id = self._active_request_message_id + self._active_request_message_id = request_message_id + try: + async with connect(url, **self._connection_kwargs()) as websocket: + self._log_ws_event( + "CONNECT_OK", + url=url, + host=self._url_host(url), + protocol=self._protocol or "-", + ) + self._emit_request(url=url, payload=payload) + request_started = time.perf_counter() + try: + await asyncio.wait_for( + websocket.send(self._serialize_message(payload)), + timeout=self._write_timeout_s, + ) + except Exception as exc: + self._log_ws_failure( + "SEND_FAIL", + exc, + url=url, + host=self._url_host(url), + **self._payload_metadata(payload), + ) + raise + reply = await self._recv_reply(websocket, timeout_s=self._read_timeout_s) + if request_message_id: + self._reply_with_message_id(reply, request_message_id) + done_fields = { + **self._payload_metadata(payload), + "duration_ms": round((time.perf_counter() - request_started) * 1000), + "stage": reply.stage, + "text_len": len(reply.text), + "text_preview": reply.text, + "done": int(reply.done), + "result_type": type(reply.result).__name__ if reply.result is not None else "-", + } + self._log_ws_event( + "TURN_DONE", + url=url, + host=self._url_host(url), + **done_fields, + ) + self._active_request_message_id = previous_request_message_id + return reply + except asyncio.TimeoutError: + self._active_request_message_id = previous_request_message_id + self._log_ws_event( + "REQUEST_TIMEOUT", + url=url, + host=self._url_host(url), + **self._payload_metadata(payload), + ) + raise + except Exception as exc: + self._active_request_message_id = previous_request_message_id + self._log_ws_failure( + "REQUEST_FAIL", + exc, + url=url, + host=self._url_host(url), + **self._payload_metadata(payload), + ) + raise + + async def prepare( + self, + elegibility: bool, + protocol: str, + ) -> None: + if self._timeline is not None: + self._timeline.emit( + "backend_prepare_started", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + ) + self._elegibility = bool(elegibility) + self._protocol = str(protocol or "") + self._last_stage = "INTRO" + self._pending_interrupt = None + self._pending_events.clear() + self._pending_ready_reply = None + if self._uses_persistent_session(): + async with self._session_lock: + await self._close_persistent_session_locked() + await self._open_persistent_session_locked(capture_ready=True) + if self._timeline is not None: + self._timeline.emit( + "backend_prepare_completed", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + protocol=self._protocol, + ) + + async def run(self, user_input: Any) -> BackendReply: + sent_payload = True + if self._uses_persistent_session(): + reply, sent_payload = await self._run_with_persistent_session(user_input) + else: + reply = await self._request(self._build_turn_payload(user_input)) + + if sent_payload: + self._clear_pending_state() + self._last_stage = reply.stage or self._last_stage or self._default_stage + return BackendReply( + stage=self._last_stage, + text=reply.text, + done=bool(reply.done) or self._last_stage == "DONE", + export_payload=reply.result, + metadata=reply.metadata, + ) + + async def set_interruption( + self, + interrupted: bool, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + if not interrupted: + self._pending_interrupt = None + if self._timeline is not None: + self._timeline.emit( + "backend_set_interruption", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + interrupted=False, + ) + return + + speech_id = self._string_or_empty(speech_id) + if not speech_id: + self._pending_interrupt = None + if self._timeline is not None: + self._timeline.emit( + "backend_set_interruption", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + interrupted=True, + skipped=True, + listened_text=listened_text, + speech_id="", + reason="missing_speech_id", + ) + return + self._pending_interrupt = { + "speech_id": speech_id, + "heard_text": listened_text or "", + } + if self._timeline is not None: + self._timeline.emit( + "backend_set_interruption", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + interrupted=True, + skipped=bool(skipped), + listened_text=listened_text, + speech_id=speech_id, + ) + + async def set_processing_interruption( + self, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + await self.set_interruption( + True, + listened_text=listened_text, + skipped=skipped, + speech_id=speech_id, + ) + if self._pending_interrupt is not None: + self._pending_interrupt["_interruption_field"] = "processing_interruption" + + async def inject_idle_nudge(self, nudge_text: str) -> None: + text = (nudge_text or "").strip() + if not text: + return + + event = { + "type": "idle_nudge", + "text": text, + } + # Cada frase de inatividade substitui a anterior: o cliente responde ao + # que ouviu por ultimo, e as intermediarias so empilham falas do agente + # no historico remoto -- inclusive o aviso de encerramento, que passa a + # constar como dito logo antes de a conversa seguir normalmente. + if self._pending_events and self._pending_events[-1].get("type") == "idle_nudge": + self._pending_events[-1] = event + else: + self._pending_events.append(event) + if self._timeline is not None: + self._timeline.emit( + "backend_idle_nudge_buffered", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + text=text, + ) + + async def end_service_once(self) -> BackendReply: + async with self._end_lock: + if self._ended: + return self._end_reply or BackendReply(stage="DONE", done=True, export_payload=[]) + + self._ended = True + if self._timeline is not None: + self._timeline.emit( + "backend_end_started", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + ) + try: + if self._uses_persistent_session(): + async with self._session_lock: + await self._open_persistent_session_locked( + capture_ready=self._session_websocket is None + ) + self._pending_ready_reply = None + if self._reader_task is not None and self._reader_task.done(): + if self._reader_error is not None: + raise RuntimeError("Persistent websocket reader failed") from self._reader_error + reply = self._terminal_reply or _RemoteAgentReply( + stage="DONE", + text="", + done=True, + result=[], + ) + elif not self._persistent_session_supports_action("end"): + reply = self._terminal_reply or _RemoteAgentReply( + stage="DONE", + text="", + done=True, + result=[], + ) + else: + try: + reply = await self._send_persistent_payload_locked(self._build_end_payload()) + except RuntimeError as exc: + if not self._is_unsupported_action_error(exc, "end"): + raise + reply = self._terminal_reply or _RemoteAgentReply( + stage="DONE", + text="", + done=True, + result=[], + ) + await self._close_persistent_session_locked() + else: + reply = await self._request(self._build_end_payload()) + except Exception: + logger.exception("[remote-agent] end_service websocket falhou") + if self._uses_persistent_session(): + async with self._session_lock: + await self._close_persistent_session_locked() + if self._timeline is not None: + self._timeline.emit( + "backend_end_failed", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + ) + self._end_reply = BackendReply(stage="DONE", done=True, export_payload=[]) + return self._end_reply + + self._clear_pending_state() + stage = (reply.stage or self._last_stage or self._default_stage).upper() + self._last_stage = stage + self._end_reply = BackendReply( + stage=stage, + text=reply.text, + done=bool(reply.done) or stage == "DONE", + export_payload=reply.result if reply.result is not None else [], + metadata=reply.metadata, + ) + if self._timeline is not None: + self._timeline.emit( + "backend_end_completed", + backend=self._backend_label, + backend_family="remote_ws", + agent=self._current_agent_name(), + result_type=type(self._end_reply.export_payload).__name__, + ) + return self._end_reply diff --git a/src/app/livekit/adapters/speech_service.py b/src/app/livekit/adapters/speech_service.py new file mode 100644 index 0000000..78c742d --- /dev/null +++ b/src/app/livekit/adapters/speech_service.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +import inspect +import re +from typing import Any + + +def normalize_text(text: str) -> str: + text = re.sub(r"\s+", " ", (text or "")).strip() + text = re.sub(r"\s+([,.;:!?…])", r"\1", text) + return text.strip() + + +def _extract_text_from_content(content: Any) -> str: + if content is None: + return "" + if isinstance(content, str): + return content + if isinstance(content, list): + parts: list[str] = [] + for item in content: + if isinstance(item, str): + parts.append(item) + else: + text = getattr(item, "text", None) + if isinstance(text, str): + parts.append(text) + else: + rendered = str(item) + if rendered and rendered != "None": + parts.append(rendered) + return "".join(parts).strip() + text = getattr(content, "text", None) + if isinstance(text, str): + return text + return str(content).strip() + + +def extract_spoken_from_speech_handle(handle: Any) -> str: + items = getattr(handle, "chat_items", None) + if not isinstance(items, list) or not items: + return "" + + for item in reversed(items): + role = getattr(item, "role", None) + if role is None or str(role).lower() == "assistant": + content = getattr(item, "content", None) + spoken = normalize_text(_extract_text_from_content(content)) + if spoken: + return spoken + + item = items[-1] + return normalize_text(_extract_text_from_content(getattr(item, "content", None))) + + +async def _maybe_await(value: Any) -> Any: + if inspect.isawaitable(value): + return await value + return value + + +async def _maybe_await_call(fn, *args, **kwargs) -> Any: + if fn is None: + return None + try: + value = fn(*args, **kwargs) + except TypeError: + value = fn(*args) + return await _maybe_await(value) + + +class SpeechService: + def __init__(self, session: Any) -> None: + self._session = session + + async def start( + self, + text: str, + *, + allow_interruptions: bool, + add_to_chat_ctx: bool = True, + audio: Any = None, + ) -> Any: + kwargs = { + "allow_interruptions": allow_interruptions, + "add_to_chat_ctx": add_to_chat_ctx, + } + if audio is not None: + kwargs["audio"] = audio + value = self._session.say(text, **kwargs) + if inspect.isawaitable(value) and not callable(getattr(value, "wait_for_playout", None)): + return await value + return value + + async def wait_for_playout(self, handle: Any) -> None: + wait_for_playout = getattr(handle, "wait_for_playout", None) + if callable(wait_for_playout): + await _maybe_await_call(wait_for_playout) + + async def interrupt(self, handle: Any, *, force: bool = False) -> None: + interrupt = getattr(handle, "interrupt", None) + if callable(interrupt): + await _maybe_await_call(interrupt, force=force) + return + + session_interrupt = getattr(self._session, "interrupt", None) + if callable(session_interrupt): + await _maybe_await_call(session_interrupt, force=force) + + def extract_spoken_text(self, handle: Any) -> str: + return extract_spoken_from_speech_handle(handle) diff --git a/src/app/livekit/adapters/xai_pool_proxy.py b/src/app/livekit/adapters/xai_pool_proxy.py new file mode 100644 index 0000000..66cac72 --- /dev/null +++ b/src/app/livekit/adapters/xai_pool_proxy.py @@ -0,0 +1,579 @@ +from __future__ import annotations + +import asyncio +import json +import logging +import os +import random +import time +from dataclasses import dataclass +from typing import Any +from urllib.parse import urlencode + +import aiohttp +from aiohttp import WSMsgType, web + +from app.livekit.adapters.xai_tts import ( + AUTH_METHOD_API_KEY, + AUTH_METHOD_CONFIG_FILE, + DEFAULT_LANGUAGE, + DEFAULT_VOICE, + OPTIMIZE_STREAMING_LATENCY, + SAMPLE_RATE, + TEXT_NORMALIZATION, + _AuthOptions, + _request_headers, + _resolve_auth_method, + _validate_config_file_auth, +) + +logger = logging.getLogger("xai_pool_proxy") + + +def _env_int(name: str, default: int) -> int: + try: + return int(os.getenv(name, str(default))) + except (TypeError, ValueError): + return default + + +def _env_float(name: str, default: float) -> float: + try: + return float(os.getenv(name, str(default))) + except (TypeError, ValueError): + return default + + +def _env_bool(name: str, default: bool) -> bool: + raw = os.getenv(name) + if raw is None: + return default + return raw.strip().lower() in {"1", "true", "yes", "on"} + + +def _auth_from_env() -> _AuthOptions: + method = _resolve_auth_method(os.getenv("XAI_TTS_AUTH_METHOD", AUTH_METHOD_API_KEY)) + api_key = os.getenv("XAI_API_KEY") + compartment_id = os.getenv("OCI_COMPARTMENT_ID") + config_file = os.getenv("OCI_CONFIG_FILE", "~/.oci/config") + profile = os.getenv("OCI_CONFIG_PROFILE", "DEFAULT") + + if method == AUTH_METHOD_API_KEY: + if not api_key: + raise RuntimeError("XAI_API_KEY is required for xAI pool API_KEY authentication") + return _AuthOptions(method=method, api_key=api_key) + + if not compartment_id: + raise RuntimeError("OCI_COMPARTMENT_ID is required for IAM xAI pool authentication") + if method == AUTH_METHOD_CONFIG_FILE: + _validate_config_file_auth(config_file, profile) + return _AuthOptions( + method=method, + compartment_id=compartment_id, + oci_config_file=config_file if method == AUTH_METHOD_CONFIG_FILE else None, + oci_profile=profile if method == AUTH_METHOD_CONFIG_FILE else None, + ) + + +@dataclass(frozen=True) +class PoolConfig: + region: str + upstream_url: str + voice: str + language: str + size: int + unavailable_free_threshold: int + recover_free_threshold: int + connect_timeout_s: float + connection_ttl_s: float + refresh_jitter_s: float + maintenance_interval_s: float + acquire_timeout_s: float + + @classmethod + def from_env(cls) -> "PoolConfig": + size = max(1, _env_int("XAI_POOL_SIZE", 50)) + unavailable = max(0, _env_int("XAI_POOL_UNAVAILABLE_FREE", 2)) + recover = max(unavailable + 1, _env_int("XAI_POOL_RECOVER_FREE", 5)) + recover = min(size, recover) + return cls( + region=(os.getenv("TIA_XAI_REGION") or "unknown").strip(), + upstream_url=(os.getenv("XAI_POOL_UPSTREAM_URL") or os.getenv("XAI_UPSTREAM_WEBSOCKET_URL") or "").strip(), + voice=(os.getenv("XAI_TTS_VOICE") or DEFAULT_VOICE).strip() or DEFAULT_VOICE, + language=(os.getenv("XAI_TTS_LANGUAGE") or DEFAULT_LANGUAGE).strip() or DEFAULT_LANGUAGE, + size=size, + unavailable_free_threshold=min(size, unavailable), + recover_free_threshold=recover, + connect_timeout_s=max(0.1, _env_float("XAI_POOL_CONNECT_TIMEOUT_S", 3.0)), + connection_ttl_s=max(30.0, _env_float("XAI_POOL_CONNECTION_TTL_S", 540.0)), + refresh_jitter_s=max(0.0, _env_float("XAI_POOL_REFRESH_JITTER_S", 45.0)), + maintenance_interval_s=max(0.5, _env_float("XAI_POOL_MAINTENANCE_INTERVAL_S", 2.0)), + acquire_timeout_s=max(0.1, _env_float("XAI_POOL_ACQUIRE_TIMEOUT_S", 2.0)), + ) + + def upstream_ws_url(self) -> str: + params = { + "voice": self.voice, + "language": self.language, + "codec": "pcm", + "sample_rate": SAMPLE_RATE, + "optimize_streaming_latency": OPTIMIZE_STREAMING_LATENCY, + "text_normalization": str(TEXT_NORMALIZATION).lower(), + } + return f"{self.upstream_url}?{urlencode(params)}" + + +class PoolUnavailable(RuntimeError): + pass + + +class UpstreamSlot: + def __init__(self, slot_id: int, *, config: PoolConfig, session: aiohttp.ClientSession, auth: _AuthOptions) -> None: + self.slot_id = slot_id + self.config = config + self.session = session + self.auth = auth + self.ws: aiohttp.ClientWebSocketResponse | None = None + self.lock = asyncio.Lock() + self.opened_at = 0.0 + self.last_used_at = 0.0 + self.refresh_deadline = 0.0 + self.connect_failures = 0 + + @property + def leased(self) -> bool: + return self.lock.locked() + + @property + def healthy(self) -> bool: + ws = self.ws + return bool(ws is not None and not ws.closed and ws.exception() is None) + + @property + def needs_refresh(self) -> bool: + return bool(self.healthy and not self.leased and self.refresh_deadline > 0 and time.monotonic() >= self.refresh_deadline) + + async def connect(self) -> None: + if self.healthy: + return + await self.close() + url = self.config.upstream_ws_url() + started = time.perf_counter() + try: + self.ws = await asyncio.wait_for( + self.session.ws_connect( + url, + headers=_request_headers(self.auth, url), + heartbeat=None, + autoclose=True, + ), + timeout=self.config.connect_timeout_s, + ) + except Exception: + self.connect_failures += 1 + raise + self.opened_at = time.monotonic() + self.last_used_at = self.opened_at + jitter = random.uniform(0.0, min(self.config.refresh_jitter_s, max(0.0, self.config.connection_ttl_s - 1.0))) + self.refresh_deadline = self.opened_at + self.config.connection_ttl_s - jitter + logger.info( + "XAI_POOL_SLOT_OPENED slot=%s region=%s connect_ms=%s refresh_in_s=%.1f", + self.slot_id, + self.config.region, + round((time.perf_counter() - started) * 1000), + max(0.0, self.refresh_deadline - time.monotonic()), + ) + + async def close(self) -> None: + ws, self.ws = self.ws, None + if ws is not None and not ws.closed: + try: + await ws.close() + except Exception: + logger.debug("failed closing xAI pool slot=%s", self.slot_id, exc_info=True) + + async def refresh(self) -> None: + if self.leased: + return + await self.close() + await self.connect() + + +class RegionalXAIPool: + def __init__(self, config: PoolConfig) -> None: + if not config.upstream_url: + raise RuntimeError("XAI_POOL_UPSTREAM_URL is required") + self.config = config + self.auth = _auth_from_env() + self.session: aiohttp.ClientSession | None = None + self.slots: list[UpstreamSlot] = [] + self._condition = asyncio.Condition() + self._maintenance_task: asyncio.Task[None] | None = None + self._draining = False + self._ready = False + self._started = False + self.total_acquires = 0 + self.total_acquire_timeouts = 0 + self.total_proxy_failures = 0 + + async def start(self) -> None: + if self._started: + return + connector = aiohttp.TCPConnector(limit=max(self.config.size * 2, 100), ttl_dns_cache=300) + self.session = aiohttp.ClientSession(connector=connector) + self.slots = [UpstreamSlot(i + 1, config=self.config, session=self.session, auth=self.auth) for i in range(self.config.size)] + # Prewarm in bounded waves so startup does not create one handshake burst. + concurrency = max(1, min(self.config.size, _env_int("XAI_POOL_PREWARM_CONCURRENCY", 5))) + sem = asyncio.Semaphore(concurrency) + + async def open_slot(slot: UpstreamSlot) -> None: + async with sem: + try: + await slot.connect() + except Exception as exc: + logger.warning("XAI_POOL_PREWARM_FAILED slot=%s error=%s", slot.slot_id, type(exc).__name__) + await asyncio.sleep(max(0.0, _env_float("XAI_POOL_PREWARM_STAGGER_S", 0.05))) + + await asyncio.gather(*(open_slot(slot) for slot in self.slots)) + self._started = True + self._recompute_ready() + self._maintenance_task = asyncio.create_task(self._maintenance_loop(), name="xai-pool-maintenance") + logger.info("XAI_POOL_STARTED region=%s size=%s healthy=%s ready=%s", self.config.region, self.config.size, self.healthy_count, self._ready) + + async def stop(self) -> None: + self._draining = True + self._ready = False + if self._maintenance_task is not None: + self._maintenance_task.cancel() + try: + await self._maintenance_task + except asyncio.CancelledError: + pass + self._maintenance_task = None + for slot in self.slots: + await slot.close() + if self.session is not None: + await self.session.close() + self.session = None + self._started = False + + @property + def healthy_count(self) -> int: + return sum(1 for slot in self.slots if slot.healthy) + + @property + def leased_count(self) -> int: + return sum(1 for slot in self.slots if slot.leased) + + @property + def free_healthy_count(self) -> int: + return sum(1 for slot in self.slots if slot.healthy and not slot.leased) + + @property + def ready(self) -> bool: + self._recompute_ready() + return self._ready + + def _recompute_ready(self) -> None: + if self._draining or not self._started: + self._ready = False + return + free = self.free_healthy_count + if self._ready: + if free <= self.config.unavailable_free_threshold: + self._ready = False + else: + if free >= self.config.recover_free_threshold: + self._ready = True + + async def set_draining(self, draining: bool = True) -> None: + self._draining = draining + self._recompute_ready() + async with self._condition: + self._condition.notify_all() + + async def acquire(self) -> UpstreamSlot: + deadline = time.monotonic() + self.config.acquire_timeout_s + while True: + if self._draining: + raise PoolUnavailable("pool is draining") + for slot in self.slots: + if not slot.healthy or slot.leased: + continue + if not slot.lock.locked(): + await slot.lock.acquire() + if not slot.healthy: + slot.lock.release() + continue + self.total_acquires += 1 + self._recompute_ready() + return slot + remaining = deadline - time.monotonic() + if remaining <= 0: + self.total_acquire_timeouts += 1 + self._recompute_ready() + raise PoolUnavailable("no free healthy xAI connection") + async with self._condition: + try: + await asyncio.wait_for(self._condition.wait(), timeout=min(remaining, 0.25)) + except asyncio.TimeoutError: + pass + + async def release(self, slot: UpstreamSlot, *, healthy: bool = True) -> None: + slot.last_used_at = time.monotonic() + if not healthy: + await slot.close() + if slot.lock.locked(): + slot.lock.release() + self._recompute_ready() + async with self._condition: + self._condition.notify_all() + + async def _maintenance_loop(self) -> None: + while True: + await asyncio.sleep(self.config.maintenance_interval_s) + for slot in self.slots: + if slot.leased: + continue + try: + if slot.needs_refresh or not slot.healthy: + await slot.refresh() + except Exception as exc: + logger.warning("XAI_POOL_SLOT_RECOVERY_FAILED slot=%s error=%s", slot.slot_id, type(exc).__name__) + self._recompute_ready() + async with self._condition: + self._condition.notify_all() + + def status(self) -> dict[str, Any]: + self._recompute_ready() + return { + "status": "ready" if self._ready else "not_ready", + "region": self.config.region, + "draining": self._draining, + "configured": self.config.size, + "healthy": self.healthy_count, + "leased": self.leased_count, + "free": self.free_healthy_count, + "unavailable_free_threshold": self.config.unavailable_free_threshold, + "recover_free_threshold": self.config.recover_free_threshold, + "total_acquires": self.total_acquires, + "total_acquire_timeouts": self.total_acquire_timeouts, + "total_proxy_failures": self.total_proxy_failures, + } + + +POOL: RegionalXAIPool | None = None + + +def _client_query_matches(request: web.Request, config: PoolConfig) -> bool: + # The pool is prewarmed for one voice/language/sample-rate profile. Explicit + # mismatch is rejected instead of silently synthesizing with the wrong voice. + expected = { + "voice": config.voice, + "language": config.language, + "codec": "pcm", + "sample_rate": str(SAMPLE_RATE), + } + for key, value in expected.items(): + incoming = request.query.get(key) + if incoming is not None and incoming != value: + return False + return True + + +async def _relay_upstream_until_boundary(client: web.WebSocketResponse, slot: UpstreamSlot) -> bool: + """Relay one provider response boundary. Return True only after audio.done.""" + ws = slot.ws + if ws is None: + return False + while True: + msg = await ws.receive() + if msg.type == WSMsgType.TEXT: + await client.send_str(msg.data) + try: + payload = json.loads(msg.data) + except Exception: + payload = {} + if payload.get("type") == "audio.done": + return True + if payload.get("type") in {"error", "response.error"}: + return False + elif msg.type == WSMsgType.BINARY: + await client.send_bytes(msg.data) + elif msg.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.ERROR}: + return False + + +async def websocket_proxy(request: web.Request) -> web.StreamResponse: + pool = POOL + if pool is None: + raise web.HTTPServiceUnavailable(text="pool not initialized") + if not _client_query_matches(request, pool.config): + raise web.HTTPBadRequest(text="voice/language/codec/sample_rate differs from prewarmed pool profile") + + client = web.WebSocketResponse(heartbeat=20.0, max_msg_size=8 * 1024 * 1024) + await client.prepare(request) + leased_slot: UpstreamSlot | None = None + slot_healthy = True + try: + async for msg in client: + if msg.type != WSMsgType.TEXT: + if msg.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.ERROR}: + break + continue + try: + payload = json.loads(msg.data) + except Exception: + await client.send_str(json.dumps({"type": "error", "message": "invalid json"})) + continue + msg_type = str(payload.get("type") or "") + + if msg_type == "text.clear": + if leased_slot is not None: + await pool.release(leased_slot, healthy=slot_healthy) + leased_slot = None + try: + leased_slot = await pool.acquire() + slot_healthy = True + except PoolUnavailable as exc: + await client.send_str(json.dumps({"type": "error", "message": str(exc), "code": "xai_pool_exhausted"})) + await client.close(code=1013, message=b"xAI pool exhausted") + break + assert leased_slot.ws is not None + try: + await leased_slot.ws.send_str(msg.data) + # text.clear has its own acknowledgement and must be forwarded + # before the client sends text.delta/text.done. + while True: + ack = await leased_slot.ws.receive() + if ack.type != WSMsgType.TEXT: + slot_healthy = False + raise RuntimeError("xAI clear acknowledgement failed") + await client.send_str(ack.data) + try: + ack_payload = json.loads(ack.data) + except Exception: + ack_payload = {} + ack_type = str(ack_payload.get("type") or "") + if ack_type == "audio.clear": + break + if ack_type in {"error", "response.error"}: + slot_healthy = False + raise RuntimeError("xAI clear returned error") + except Exception: + slot_healthy = False + pool.total_proxy_failures += 1 + await client.close(code=1011, message=b"xAI upstream clear failed") + break + continue + + if leased_slot is None: + await client.send_str(json.dumps({"type": "error", "message": "text.clear required before synthesis"})) + continue + + assert leased_slot.ws is not None + try: + await leased_slot.ws.send_str(msg.data) + if msg_type == "text.done": + slot_healthy = await _relay_upstream_until_boundary(client, leased_slot) + await pool.release(leased_slot, healthy=slot_healthy) + leased_slot = None + if not slot_healthy: + pool.total_proxy_failures += 1 + await client.close(code=1011, message=b"xAI upstream failed") + break + except Exception: + slot_healthy = False + pool.total_proxy_failures += 1 + await client.close(code=1011, message=b"xAI upstream failure") + break + finally: + if leased_slot is not None: + await pool.release(leased_slot, healthy=False) + return client + + +async def healthz(_: web.Request) -> web.Response: + pool = POOL + if pool is None or not pool._started: + return web.json_response({"status": "starting"}, status=503) + return web.json_response({"status": "ok", "pool": pool.status()}) + + +async def readyz(_: web.Request) -> web.Response: + pool = POOL + if pool is None: + return web.json_response({"status": "not_ready", "reason": "not_initialized"}, status=503) + payload = pool.status() + return web.json_response(payload, status=200 if pool.ready else 503) + + +async def status(_: web.Request) -> web.Response: + pool = POOL + return web.json_response(pool.status() if pool is not None else {"status": "not_initialized"}) + + +async def drain(_: web.Request) -> web.Response: + pool = POOL + if pool is not None: + await pool.set_draining(True) + return web.json_response({"status": "draining"}) + + +async def metrics(_: web.Request) -> web.Response: + pool = POOL + values = pool.status() if pool is not None else {} + region = str(values.get("region", "unknown")).replace('"', "") + lines = [ + "# TYPE tia_xai_pool_connections gauge", + f'tia_xai_pool_connections{{region="{region}",state="healthy"}} {values.get("healthy", 0)}', + f'tia_xai_pool_connections{{region="{region}",state="leased"}} {values.get("leased", 0)}', + f'tia_xai_pool_connections{{region="{region}",state="free"}} {values.get("free", 0)}', + "# TYPE tia_xai_pool_acquires_total counter", + f'tia_xai_pool_acquires_total{{region="{region}"}} {values.get("total_acquires", 0)}', + "# TYPE tia_xai_pool_acquire_timeouts_total counter", + f'tia_xai_pool_acquire_timeouts_total{{region="{region}"}} {values.get("total_acquire_timeouts", 0)}', + "# TYPE tia_xai_pool_proxy_failures_total counter", + f'tia_xai_pool_proxy_failures_total{{region="{region}"}} {values.get("total_proxy_failures", 0)}', + ] + return web.Response(text="\n".join(lines) + "\n", content_type="text/plain") + + +async def on_startup(app: web.Application) -> None: + global POOL + POOL = RegionalXAIPool(PoolConfig.from_env()) + await POOL.start() + + +async def on_cleanup(app: web.Application) -> None: + global POOL + if POOL is not None: + await POOL.stop() + POOL = None + + +def create_app() -> web.Application: + app = web.Application() + app.router.add_get("/xai/v1/tts", websocket_proxy) + app.router.add_get("/healthz", healthz) + app.router.add_get("/readyz", readyz) + app.router.add_get("/pool/status", status) + app.router.add_post("/drain", drain) + app.router.add_get("/metrics", metrics) + app.on_startup.append(on_startup) + app.on_cleanup.append(on_cleanup) + return app + + +def main() -> None: + logging.basicConfig(level=os.getenv("LOG_LEVEL", "INFO")) + web.run_app( + create_app(), + host=os.getenv("XAI_POOL_BIND_HOST", "0.0.0.0"), + port=_env_int("XAI_POOL_PORT", 18100), + access_log=logger if _env_bool("XAI_POOL_ACCESS_LOG", False) else None, + ) + + +if __name__ == "__main__": + main() diff --git a/src/app/livekit/adapters/xai_tts.py b/src/app/livekit/adapters/xai_tts.py new file mode 100644 index 0000000..85c9ad9 --- /dev/null +++ b/src/app/livekit/adapters/xai_tts.py @@ -0,0 +1,1519 @@ +from __future__ import annotations + +import asyncio +import base64 +import json +import logging +import os +import time +import unicodedata +import wave +import weakref +from collections.abc import AsyncIterable +from dataclasses import dataclass, replace +from typing import Literal, cast +from urllib.parse import urlencode + +import aiohttp +import oci +import requests +from livekit.agents import ( + APIConnectionError, + APIConnectOptions, + APIStatusError, + APITimeoutError, + tts, + utils, +) +from livekit.agents.metrics import TTSMetrics +from livekit.agents.metrics.base import Metadata +from livekit.agents.telemetry import trace_types +from livekit.agents.types import DEFAULT_API_CONNECT_OPTIONS, NOT_GIVEN, NotGivenOr +from livekit.agents.utils import is_given + +from app.livekit.adapters.audio_gain import GainEmitter, tts_output_gain_from_env +from app.livekit.runtime.initial_greeting_audio_cache import ( + InitialGreetingAudioCache, + InitialGreetingCacheKey, +) + +logger = logging.getLogger(__name__) + + +SAMPLE_RATE = 24000 +NUM_CHANNELS = 1 +PCM_BYTES_PER_SAMPLE = 2 + +DEFAULT_XAI_WEBSOCKET_URL = ( + "wss://inference.generativeai.us-chicago-1.oci.oraclecloud.com/xai/v1/tts" +) +DEFAULT_VOICE = "c8x2ieiocufs" +DEFAULT_LANGUAGE = "pt-BR" +TTS_EMPTY_FRAME_RETRY_TIMEOUT_ENV = "TTS_EMPTY_FRAME_RETRY_TIMEOUT_S" +TTS_FIRST_FRAME_TIMEOUT_ENV = "TTS_FIRST_FRAME_TIMEOUT_S" +TTS_UNDERFLOW_ERROR_ENV = "TTS_UNDERFLOW_ERROR_MS" +TTS_TOTAL_TIMEOUT_ENV = "TTS_TOTAL_TIMEOUT_S" +_AUDIO_EMITTER_FRAME_SIZE_MS = 40 +_ALLOWED_TTS_PUNCTUATION = frozenset(".,;:?!-()[]{}\"'") | { + "\u2013", + "\u2014", + "\u2018", + "\u2019", + "\u201c", + "\u201d", + "\u2026", +} +OPTIMIZE_STREAMING_LATENCY = 1 +TEXT_NORMALIZATION = True +WEBSOCKET_HEARTBEAT: float | None = None +# Reconnect before the service can silently expire an idle socket, which also +# refreshes IAM authorization at the next WebSocket handshake. +WEBSOCKET_IDLE_TTL_SECONDS = 290.0 + +AUTH_METHOD_API_KEY = "API_KEY" +AUTH_METHOD_CONFIG_FILE = "CONFIG_FILE" +AUTH_METHOD_INSTANCE_PRINCIPAL = "INSTANCE_PRINCIPAL" +AUTH_METHOD_OKE_WORKLOAD_IDENTITY = "OKE_WORKLOAD_IDENTITY" +AuthMethod = Literal[ + "API_KEY", + "CONFIG_FILE", + "INSTANCE_PRINCIPAL", + "OKE_WORKLOAD_IDENTITY", +] +_IAM_AUTH_METHODS = { + AUTH_METHOD_CONFIG_FILE, + AUTH_METHOD_INSTANCE_PRINCIPAL, + AUTH_METHOD_OKE_WORKLOAD_IDENTITY, +} +_AUTH_METHODS = _IAM_AUTH_METHODS | {AUTH_METHOD_API_KEY} + + +def _runtime_logger() -> logging.Logger: + app_logger = logging.getLogger(os.getenv("APP_LOGGER_NAME", "agent_internal_stt")) + return app_logger if app_logger.handlers else logger + + +def _elapsed_ms(started_at: float) -> int: + if started_at <= 0.0: + return 0 + return max(0, round((time.perf_counter() - started_at) * 1000)) + + +def _env_timeout_s(env_name: str, default: float) -> float: + try: + timeout_s = float(str(os.getenv(env_name, str(default)) or str(default))) + except ValueError: + timeout_s = default + return max(0.0, timeout_s) + + +def _first_frame_timeout_s() -> float: + if os.getenv(TTS_FIRST_FRAME_TIMEOUT_ENV) is not None: + return _env_timeout_s(TTS_FIRST_FRAME_TIMEOUT_ENV, 3.0) + return _env_timeout_s(TTS_EMPTY_FRAME_RETRY_TIMEOUT_ENV, 3.0) + + +def _underflow_error_ms() -> int: + """Maximum estimated PCM debt tolerated after playout has started.""" + try: + value = int(str(os.getenv(TTS_UNDERFLOW_ERROR_ENV, "1000") or "1000")) + except ValueError: + value = 1000 + return max(0, value) + + +def _turn_total_timeout_s() -> float: + """Absolute ceiling from ``text.delta`` to ``audio.done`` for one turn.""" + return _env_timeout_s(TTS_TOTAL_TIMEOUT_ENV, 60.0) + + +def _pcm_duration_s(payload_bytes: int) -> float: + return max(0, payload_bytes) / (SAMPLE_RATE * PCM_BYTES_PER_SAMPLE) + + +def _estimated_pcm_balance_s( + *, + pcm_duration_s: float, + first_pcm_released_at: float | None, + now: float, +) -> float: + """PCM received minus playout elapsed; no implicit LiveKit prebuffer.""" + if first_pcm_released_at is None: + return 0.0 + return pcm_duration_s - max(0.0, now - first_pcm_released_at) + + +def _sanitize_tts_text(text: str) -> str: + raw_text = str(text or "") + parts: list[str] = [] + for index, char in enumerate(raw_text): + category = unicodedata.category(char) + if ( + category[0] in {"L", "N"} + or char in _ALLOWED_TTS_PUNCTUATION + or ( + char == "/" + and index > 0 + and index + 1 < len(raw_text) + and raw_text[index - 1].isdigit() + and raw_text[index + 1].isdigit() + ) + ): + parts.append(char) + elif ( + char == "/" + and index > 0 + and index + 1 < len(raw_text) + and raw_text[index - 1].isalpha() + and raw_text[index + 1].isalpha() + ): + parts.append(" ou ") + elif char.isspace(): + parts.append(" ") + else: + parts.append(" ") + return " ".join("".join(parts).split()) + + +@dataclass +class _TTSOptions: + base_url: str + voice: str + language: str + + +@dataclass(frozen=True) +class _AuthOptions: + method: AuthMethod + api_key: str | None = None + compartment_id: str | None = None + oci_config_file: str | None = None + oci_profile: str | None = None + + +@dataclass +class _TurnTiming: + """Internal timing data for one xAI provider turn.""" + + turn_index: int + clear_rtt_ms: float | None + provider_ttfb_ms: float | None + end_to_end_ttfb_ms: float | None + max_audio_delta_gap_ms: int + max_playout_underrun_0ms: int + connection_reused: bool + reconnected: bool + discarded_messages: list[str] + provider_synthesis_ms: float | None = None + xai_underrun_estimado_ms: int = 0 + xai_micro_underflows: int = 0 + xai_avg_underrun_ms: int = 0 + underflow_error_ms: int = 1000 + pcm_bytes: int = 0 + pcm_duration_ms: int = 0 + empty_audio_deltas: int = 0 + attempts: int = 1 + connection_queue_wait_ms: int = 0 + turn_gap_after_previous_ms: int | None = None + clear_discarded_message_count: int = 0 + clear_discarded_audio_bytes: int = 0 + + +@dataclass +class _TurnResult: + audio_emitted: bool + trace_id: str | None = None + provider_ttfb: float = -1.0 + timing: _TurnTiming | None = None + + +class _XAIConnectionClosed(Exception): + def __init__(self, message: str, *, audio_emitted: bool) -> None: + super().__init__(message) + self.audio_emitted = audio_emitted + + +class _XAIEmptyAudioDone(Exception): + """Clean provider boundary with no non-empty PCM; retry once on this socket.""" + + +class _XAIPartialAudioFailure(Exception): + """A response already reached playout and therefore must never be replayed.""" + + +def _resolve_auth_method(value: str) -> AuthMethod: + resolved = value.upper() + if resolved not in _AUTH_METHODS: + raise ValueError( + "xAI TTS auth_method must be one of: " + ", ".join(sorted(_AUTH_METHODS)) + ) + return cast(AuthMethod, resolved) + + +def _validate_config_file_auth(config_file: str, profile: str) -> None: + try: + oci.config.from_file(os.path.expanduser(config_file), profile) + except Exception as exc: + raise ValueError( + "CONFIG_FILE authentication requires a readable OCI config file and " + f"profile (resolved path: {os.path.expanduser(config_file)!r}, " + f"profile: {profile!r})" + ) from exc + + +def _make_iam_signer(auth: _AuthOptions): + if auth.method == AUTH_METHOD_CONFIG_FILE: + assert auth.oci_config_file is not None + assert auth.oci_profile is not None + oci_config = oci.config.from_file( + os.path.expanduser(auth.oci_config_file), auth.oci_profile + ) + if "security_token_file" not in oci_config: + return oci.signer.Signer.from_config(oci_config) + + private_key = oci.signer.load_private_key_from_file(oci_config["key_file"]) + with open(oci_config["security_token_file"], encoding="utf-8") as token_file: + token = token_file.read().strip() + return oci.auth.signers.SecurityTokenSigner(token, private_key) + + if auth.method == AUTH_METHOD_INSTANCE_PRINCIPAL: + return oci.auth.signers.InstancePrincipalsSecurityTokenSigner() + + if auth.method == AUTH_METHOD_OKE_WORKLOAD_IDENTITY: + return oci.auth.signers.get_oke_workload_identity_resource_principal_signer() + + raise ValueError(f"unsupported IAM authentication method: {auth.method}") + + +def _request_headers(auth: _AuthOptions, uri: str) -> dict[str, str]: + if auth.method == AUTH_METHOD_API_KEY: + assert auth.api_key is not None + return {"Authorization": f"Bearer {auth.api_key}"} + + signer = _make_iam_signer(auth) + https_uri = uri.replace("wss://", "https://", 1) + prepared = requests.Request("GET", https_uri).prepare() + signer.do_request_sign(prepared) + headers = dict(prepared.headers) + assert auth.compartment_id is not None + headers["opc-compartment-id"] = auth.compartment_id + return headers + + +class OraclexAITTS(tts.TTS): + def __init__( + self, + *, + api_key: NotGivenOr[str] = NOT_GIVEN, + auth_method: NotGivenOr[AuthMethod] = NOT_GIVEN, + compartment_id: NotGivenOr[str] = NOT_GIVEN, + oci_config_file: NotGivenOr[str] = NOT_GIVEN, + oci_profile: NotGivenOr[str] = NOT_GIVEN, + base_url: NotGivenOr[str] = NOT_GIVEN, + websocket_url: str | None = None, + voice: str = DEFAULT_VOICE, + language: str = DEFAULT_LANGUAGE, + http_session: aiohttp.ClientSession | None = None, + initial_greeting_audio_cache: InitialGreetingAudioCache | None = None, + initial_greeting_agent: str = "", + ) -> None: + """ + Create a new instance of the xAI TTS. + + Args: + voice (str, optional): The voice ID for the desired voice. + language (str, optional): Language code for synthesis. + api_key (str | None, optional): API key used with auth_method=API_KEY. Defaults to XAI_API_KEY. + auth_method (str, optional): API_KEY, CONFIG_FILE, INSTANCE_PRINCIPAL, or OKE_WORKLOAD_IDENTITY. Defaults to XAI_TTS_AUTH_METHOD or API_KEY. + compartment_id (str, optional): OCI compartment OCID required for IAM authentication. Defaults to OCI_COMPARTMENT_ID. + oci_config_file (str, optional): OCI config file for CONFIG_FILE. Defaults to OCI_CONFIG_FILE or ~/.oci/config. + oci_profile (str, optional): OCI profile for CONFIG_FILE. Defaults to OCI_CONFIG_PROFILE or DEFAULT. + base_url (str, optional): WebSocket base URL for xAI TTS. Precedence: explicit argument, XAI_WEBSOCKET_URL environment variable, then the built-in default. + http_session (aiohttp.ClientSession | None, optional): An existing aiohttp ClientSession to use. + """ + super().__init__( + capabilities=tts.TTSCapabilities(streaming=True), + sample_rate=SAMPLE_RATE, + num_channels=NUM_CHANNELS, + ) + + resolved_method = _resolve_auth_method( + str(auth_method) + if is_given(auth_method) + else os.environ.get("XAI_TTS_AUTH_METHOD", AUTH_METHOD_API_KEY) + ) + + resolved_key: str | None = ( + str(api_key) if is_given(api_key) else os.environ.get("XAI_API_KEY") + ) + resolved_compartment_id = ( + str(compartment_id) + if is_given(compartment_id) + else os.environ.get("OCI_COMPARTMENT_ID") + ) + resolved_config_file = ( + str(oci_config_file) + if is_given(oci_config_file) + else os.environ.get("OCI_CONFIG_FILE", "~/.oci/config") + ) + resolved_profile = ( + str(oci_profile) + if is_given(oci_profile) + else os.environ.get("OCI_CONFIG_PROFILE", "DEFAULT") + ) + + if resolved_method == AUTH_METHOD_API_KEY: + if not resolved_key: + raise ValueError( + "xAI API key is required for API_KEY authentication, either as " + "argument or set XAI_API_KEY environment variable" + ) + else: + if is_given(api_key): + raise ValueError("api_key is only valid with auth_method=API_KEY") + if not resolved_compartment_id: + raise ValueError( + "compartment_id is required for IAM authentication; pass it as " + "an argument or set OCI_COMPARTMENT_ID" + ) + if resolved_method == AUTH_METHOD_CONFIG_FILE: + _validate_config_file_auth(resolved_config_file, resolved_profile) + + self._auth = _AuthOptions( + method=resolved_method, + api_key=resolved_key if resolved_method == AUTH_METHOD_API_KEY else None, + compartment_id=( + resolved_compartment_id + if resolved_method in _IAM_AUTH_METHODS + else None + ), + oci_config_file=( + resolved_config_file + if resolved_method == AUTH_METHOD_CONFIG_FILE + else None + ), + oci_profile=( + resolved_profile if resolved_method == AUTH_METHOD_CONFIG_FILE else None + ), + ) + if websocket_url is not None: + if is_given(base_url): + raise ValueError("base_url and websocket_url cannot be used together") + resolved_base_url = str(websocket_url).strip() + if not resolved_base_url: + raise ValueError("xAI TTS websocket_url must be a non-empty string") + elif is_given(base_url): + resolved_base_url = str(base_url).strip() + if not resolved_base_url: + raise ValueError("xAI TTS base_url must be a non-empty string") + else: + resolved_base_url = ( + os.environ.get("XAI_WEBSOCKET_URL") or DEFAULT_XAI_WEBSOCKET_URL + ) + self._opts = _TTSOptions( + base_url=resolved_base_url, + voice=(voice or DEFAULT_VOICE).strip() or DEFAULT_VOICE, + language=(language or DEFAULT_LANGUAGE).strip() or DEFAULT_LANGUAGE, + ) + + self._session = http_session + self._streams = weakref.WeakSet[SynthesizeStream]() + self._connection_lock = asyncio.Lock() + self._connection: _Connection | None = None + self._retired_connections: set[_Connection] = set() + self._prewarm_task: asyncio.Task[None] | None = None + self._output_gain = tts_output_gain_from_env() + self._initial_greeting_audio_cache = initial_greeting_audio_cache + self._initial_greeting_agent = ( + (initial_greeting_agent or "unknown").strip().lower() + ) + self._initial_greeting_capture_key: InitialGreetingCacheKey | None = None + + def _wrap_output_gain(self, emitter: tts.AudioEmitter) -> tts.AudioEmitter: + """Aplica o ganho/limiter configurado a cada frame de saída.""" + if self._output_gain.enabled: + return GainEmitter(emitter, self._output_gain) + return emitter + + @property + def model(self) -> str: + return self._opts.voice + + @property + def provider(self) -> str: + return "xAI" + + def initial_greeting_audio(self, text: str): + """Return cached first-turn audio, or arm capture for the exact text.""" + cache = self._initial_greeting_audio_cache + if cache is None: + return None + key = cache.key_for( + agent=self._initial_greeting_agent, + text=_sanitize_tts_text(text), + provider=self.provider, + voice=self._opts.voice, + language=self._opts.language, + sample_rate=SAMPLE_RATE, + ) + if key is None: + return None + if cache.has(key): + try: + frames = cache.frames(key) + except (OSError, EOFError, ValueError, wave.Error): + cache.discard(key) + _runtime_logger().warning( + "INITIAL_GREETING_AUDIO_CACHE_FALLBACK | key=%s | reason=read_failed", + key.digest, + exc_info=True, + ) + else: + _runtime_logger().info( + "INITIAL_GREETING_AUDIO_CACHE_HIT | key=%s", key.digest + ) + return frames + self._initial_greeting_capture_key = key + _runtime_logger().info( + "INITIAL_GREETING_AUDIO_CACHE_ARMED | key=%s", key.digest + ) + return None + + def _take_initial_greeting_capture( + self, text: str + ) -> InitialGreetingCacheKey | None: + key = self._initial_greeting_capture_key + self._initial_greeting_capture_key = None + normalized = " ".join(_sanitize_tts_text(text).split()) + if key is None or key.text != normalized: + return None + return key + + def _ensure_session(self) -> aiohttp.ClientSession: + if not self._session: + self._session = utils.http_context.http_session() + return self._session + + def prewarm(self) -> None: + if self._prewarm_task is not None and not self._prewarm_task.done(): + return + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + return + + self._prewarm_task = loop.create_task(self._prewarm_connection()) + + async def connect( + self, timeout: float = DEFAULT_API_CONNECT_OPTIONS.timeout + ) -> None: + """Open (or validate) the reusable WebSocket connection. + + This is intentionally separate from :meth:`prewarm`: callers such as + readiness checks need an awaitable, deterministic connection result. + """ + await self._current_connection(timeout) + + async def prewarm_connection(self, *, timeout: float = 3.0) -> None: + await self.connect(timeout) + + async def _prewarm_connection(self) -> None: + try: + await self._current_connection(DEFAULT_API_CONNECT_OPTIONS.timeout) + except asyncio.CancelledError: + raise + except Exception: + _runtime_logger().debug( + "failed to prewarm xAI TTS websocket", exc_info=True + ) + + async def _current_connection( + self, timeout: float + ) -> tuple[_Connection, float, bool]: + async with self._connection_lock: + if self._connection and self._connection.is_usable: + return self._connection, 0.0, True + + if self._connection is not None: + stale_conn = self._connection + self._connection = None + stale_conn.mark_non_current() + self._retired_connections.add(stale_conn) + if not stale_conn.active: + await stale_conn.aclose() + self._retired_connections.discard(stale_conn) + + session = self._ensure_session() + conn = _Connection( + opts=replace(self._opts), + auth=self._auth, + session=session, + owner=self, + ) + t0 = time.perf_counter() + await conn.connect(timeout=timeout) + acquire_time = time.perf_counter() - t0 + self._connection = conn + return conn, acquire_time, False + + def _retire_current_connection(self) -> None: + if not self._connection: + return + + conn = self._connection + self._connection = None + conn.mark_non_current() + self._retired_connections.add(conn) + + if not conn.active: + try: + loop = asyncio.get_running_loop() + except RuntimeError: + return + + task = loop.create_task(conn.aclose()) + task.add_done_callback(lambda _: self._retired_connections.discard(conn)) + + def update_options( + self, + *, + voice: str | None = None, + language: str | None = None, + ) -> None: + """Update the xAI TTS configuration options.""" + changed = False + if voice and voice != self._opts.voice: + self._opts.voice = voice + changed = True + if language and language != self._opts.language: + self._opts.language = language + changed = True + + if changed: + self._retire_current_connection() + + def synthesize( + self, + text: str, + *, + conn_options: APIConnectOptions = DEFAULT_API_CONNECT_OPTIONS, + ) -> tts.ChunkedStream: + return _ChunkedStreamFromStream( + tts=self, + input_text=text, + conn_options=conn_options, + ) + + def stream( + self, *, conn_options: APIConnectOptions = DEFAULT_API_CONNECT_OPTIONS + ) -> SynthesizeStream: + # The adapter retains the original text and is the only retry owner. + stream = SynthesizeStream( + tts=self, + conn_options=APIConnectOptions( + max_retry=0, + retry_interval=conn_options.retry_interval, + timeout=conn_options.timeout, + ), + ) + self._streams.add(stream) + return stream + + async def aclose(self) -> None: + if self._prewarm_task is not None: + await utils.aio.gracefully_cancel(self._prewarm_task) + self._prewarm_task = None + + for stream in list(self._streams): + await stream.aclose() + self._streams.clear() + + connections = list(self._retired_connections) + if self._connection: + connections.append(self._connection) + self._connection = None + + for conn in connections: + await conn.aclose() + self._retired_connections.clear() + + +class _ChunkedStreamFromStream(tts.ChunkedStream): + async def _run(self, output_emitter: tts.AudioEmitter) -> None: + output_emitter.initialize( + request_id=utils.shortuuid(), + sample_rate=self._tts.sample_rate, + num_channels=self._tts.num_channels, + mime_type="audio/pcm", + frame_size_ms=_AUDIO_EMITTER_FRAME_SIZE_MS, + ) + async with self._tts.stream( + conn_options=APIConnectOptions( + max_retry=0, + timeout=self._conn_options.timeout, + ) + ) as stream: + stream.push_text(self.input_text) + stream.end_input() + async for event in stream: + output_emitter.push(event.frame.data.tobytes()) + output_emitter.flush() + + +class SynthesizeStream(tts.SynthesizeStream): + """Stream-based text-to-speech synthesis using xAI's WebSocket API.""" + + def __init__(self, *, tts: OraclexAITTS, conn_options: APIConnectOptions): + super().__init__(tts=tts, conn_options=conn_options) + self._tts: OraclexAITTS = tts + self._segment_started = False + self._segment_id = "" + self._segment_provider_ttfb = -1.0 + # Internal benchmark data; not part of LiveKit public metrics. + self._turn_timings: list[_TurnTiming] = [] + + async def _run(self, output_emitter: tts.AudioEmitter) -> None: + self._segment_id = utils.shortuuid() + output_emitter.initialize( + request_id=self._segment_id, + sample_rate=SAMPLE_RATE, + num_channels=NUM_CHANNELS, + stream=True, + mime_type="audio/pcm", + frame_size_ms=_AUDIO_EMITTER_FRAME_SIZE_MS, + ) + output_emitter = self._tts._wrap_output_gain(output_emitter) + + try: + await self._run_complete_input(output_emitter) + + if self._segment_started: + output_emitter.end_segment() + except asyncio.TimeoutError: + raise APITimeoutError() from None + except aiohttp.ClientResponseError as e: + raise APIStatusError( + message=e.message, + status_code=e.status, + request_id=self._segment_id, + body=None, + ) from None + + def _note_provider_ttfb(self, provider_ttfb: float) -> None: + if provider_ttfb >= 0.0 and self._segment_provider_ttfb < 0.0: + self._segment_provider_ttfb = provider_ttfb + + async def _metrics_monitor_task( + self, + event_aiter: AsyncIterable[tts.SynthesizedAudio], + ) -> None: + audio_duration = 0.0 + ttfb = -1.0 + request_id = "" + segment_id = "" + + def _emit_metrics(*, cancelled: bool = False) -> None: + nonlocal audio_duration, ttfb, request_id, segment_id + + if not self._started_time or self._current_attempt_has_error: + return + + duration = time.perf_counter() - self._started_time + + if not self._mtc_pending_texts: + return + + text = self._mtc_pending_texts.pop(0) + if not text: + return + + metrics = TTSMetrics( + timestamp=time.time(), + request_id=request_id, + segment_id=segment_id, + ttfb=ttfb, + duration=duration, + characters_count=len(text), + audio_duration=audio_duration, + # A barge-in cancels the synthesis task before the provider can + # produce its final frame. Preserve the partial timings/audio + # accumulated so far and make that distinction explicit. + cancelled=cancelled, + label=self._tts._label, + streamed=True, + metadata=Metadata( + model_name=self._tts.model, + model_provider=self._tts.provider, + ), + ) + if self._tts_request_span: + self._tts_request_span.set_attribute( + trace_types.ATTR_TTS_METRICS, + metrics.model_dump_json(), + ) + self._tts.emit("metrics_collected", metrics) + + audio_duration = 0.0 + ttfb = -1.0 + request_id = "" + self._started_time = 0 + self._segment_provider_ttfb = -1.0 + + async for ev in event_aiter: + if ttfb == -1.0: + if self._segment_provider_ttfb >= 0.0: + ttfb = self._segment_provider_ttfb + else: + ttfb = time.perf_counter() - self._started_time + + audio_duration += ev.frame.duration + request_id = ev.request_id + segment_id = ev.segment_id + + if ev.is_final: + _emit_metrics() + + # The event channel is closed when the synthesis task ends. If it was + # cancelled before an is_final frame (for example, by a barge-in), emit + # one partial metric for the current segment. _emit_metrics resets + # _started_time, so a segment that already emitted normally is never + # emitted a second time here. + if self._task.cancelled(): + _emit_metrics(cancelled=True) + + async def _run_complete_input(self, output_emitter: tts.AudioEmitter) -> None: + chunks: list[str] = [] + async for data in self._input_ch: + if isinstance(data, str): + chunks.append(data) + + text = "".join(chunks) + await self._synthesize_text(text, output_emitter) + + async def _synthesize_text( + self, text: str, output_emitter: tts.AudioEmitter + ) -> None: + normalized_text = _sanitize_tts_text(text) + if not normalized_text: + return + + if not self._segment_started: + output_emitter.start_segment(segment_id=self._segment_id) + self._segment_started = True + + ( + conn, + _, + connection_reused, + ) = await self._tts._current_connection(self._conn_options.timeout) + result = await conn.synthesize_turn( + normalized_text, + output_emitter=output_emitter, + stream=self, + timeout=self._conn_options.timeout, + turn_index=len(self._turn_timings), + connection_reused=connection_reused, + ) + if result.timing is not None: + self._turn_timings.append(result.timing) + self._emit_turn_timing(result.timing) + if result.provider_ttfb >= 0.0: + self._note_provider_ttfb(result.provider_ttfb) + + def _emit_turn_timing(self, timing: _TurnTiming) -> None: + self._tts.emit( + "xai_tts_turn_timing", + { + "segment_id": self._segment_id, + "provider_ttfb_ms": timing.provider_ttfb_ms, + "end_to_end_ttfb_ms": timing.end_to_end_ttfb_ms, + "max_audio_delta_gap_ms": timing.max_audio_delta_gap_ms, + "max_playout_underrun_0ms": timing.max_playout_underrun_0ms, + "provider_synthesis_ms": timing.provider_synthesis_ms, + "xai_underrun_estimado_ms": timing.xai_underrun_estimado_ms, + "xai_micro_underflows": timing.xai_micro_underflows, + "xai_avg_underrun_ms": timing.xai_avg_underrun_ms, + "underflow_error_ms": timing.underflow_error_ms, + "pcm_bytes": timing.pcm_bytes, + "pcm_duration_ms": timing.pcm_duration_ms, + "empty_audio_deltas": timing.empty_audio_deltas, + "attempts": timing.attempts, + "clear_rtt_ms": timing.clear_rtt_ms, + "connection_queue_wait_ms": timing.connection_queue_wait_ms, + "turn_gap_after_previous_ms": timing.turn_gap_after_previous_ms, + "clear_discarded_message_count": timing.clear_discarded_message_count, + "clear_discarded_audio_bytes": timing.clear_discarded_audio_bytes, + "connection_reused": timing.connection_reused, + "reconnected": timing.reconnected, + "discarded_message_count": len(timing.discarded_messages), + }, + ) + + +class _Connection: + """A single reusable xAI WebSocket. xAI TTS turns are serialized FIFO.""" + + def __init__( + self, + *, + opts: _TTSOptions, + auth: _AuthOptions, + session: aiohttp.ClientSession, + owner: OraclexAITTS | None = None, + ) -> None: + self._opts = opts + self._auth = auth + self._session = session + self._owner = owner + self._ws: aiohttp.ClientWebSocketResponse | None = None + self._is_current = True + self._closed = False + self._turn_lock = asyncio.Lock() + self._connection_id = f"xai-ws-{utils.shortuuid()}" + self._opened_at: float | None = None + self._last_activity_at: float | None = None + self._last_turn_finished_at: float | None = None + self._clear_sequence = 0 + + @property + def active(self) -> bool: + return self._turn_lock.locked() + + @property + def idle_seconds(self) -> float: + if self._last_activity_at is None: + return 0.0 + return time.perf_counter() - self._last_activity_at + + @property + def idle_expired(self) -> bool: + return ( + WEBSOCKET_IDLE_TTL_SECONDS > 0 + and self._last_activity_at is not None + and self.idle_seconds >= WEBSOCKET_IDLE_TTL_SECONDS + ) + + @property + def is_usable(self) -> bool: + ws = self._ws + return ( + self._is_current + and not self._closed + and ws is not None + and not ws.closed + and ws.exception() is None + and not self.idle_expired + ) + + def mark_non_current(self) -> None: + self._is_current = False + + def _note_activity(self) -> None: + self._last_activity_at = time.perf_counter() + + async def connect(self, *, timeout: float) -> None: + if self._closed: + raise APIConnectionError("xAI TTS connection is closed") + if self._ws is not None and not self._ws.closed: + return + + params = { + "voice": self._opts.voice, + "language": self._opts.language, + "codec": "pcm", + "sample_rate": SAMPLE_RATE, + "optimize_streaming_latency": OPTIMIZE_STREAMING_LATENCY, + "text_normalization": str(TEXT_NORMALIZATION).lower(), + } + url = f"{self._opts.base_url}?{urlencode(params)}" + connect_started_at = time.perf_counter() + _runtime_logger().info( + "XAI_TTS_WS_OPENING | connection_id=%s | voice=%s | language=%s | timeout_ms=%s", + self._connection_id, + self._opts.voice, + self._opts.language, + round(timeout * 1000), + ) + try: + self._ws = await asyncio.wait_for( + self._session.ws_connect( + url, + headers=_request_headers(self._auth, url), + heartbeat=WEBSOCKET_HEARTBEAT, + ), + timeout, + ) + self._opened_at = time.perf_counter() + self._note_activity() + _runtime_logger().info( + "XAI_TTS_WS_OPENED | connection_id=%s | connect_ms=%s", + self._connection_id, + round((self._opened_at - connect_started_at) * 1000), + ) + except ( + aiohttp.ClientConnectorError, + aiohttp.ClientConnectionResetError, + asyncio.TimeoutError, + ) as e: + # The original text is still owned by _Connection at this point. + raise _XAIConnectionClosed( + "failed to connect to xAI TTS", audio_emitted=False + ) from e + + async def synthesize_turn( + self, + text: str, + *, + output_emitter: tts.AudioEmitter, + stream: SynthesizeStream, + timeout: float, + turn_index: int, + connection_reused: bool, + ) -> _TurnResult: + queued_at = time.perf_counter() + async with self._turn_lock: + connection_queue_wait_ms = round((time.perf_counter() - queued_at) * 1000) + turn_gap_after_previous_ms = ( + round((time.perf_counter() - self._last_turn_finished_at) * 1000) + if self._last_turn_finished_at is not None + else None + ) + try: + return await self._synthesize_turn_with_retry( + text, + output_emitter=output_emitter, + stream=stream, + timeout=timeout, + turn_index=turn_index, + connection_reused=connection_reused, + connection_queue_wait_ms=connection_queue_wait_ms, + turn_gap_after_previous_ms=turn_gap_after_previous_ms, + ) + finally: + self._last_turn_finished_at = time.perf_counter() + if not self._is_current: + await self.aclose() + + async def _synthesize_turn_with_retry( + self, + text: str, + *, + output_emitter: tts.AudioEmitter, + stream: SynthesizeStream, + timeout: float, + turn_index: int, + connection_reused: bool, + connection_queue_wait_ms: int, + turn_gap_after_previous_ms: int | None, + ) -> _TurnResult: + logical_turn_start = time.perf_counter() + reconnected = False + last_error: Exception | None = None + for attempt in range(2): + try: + await self.connect(timeout=timeout) + return await self._synthesize_turn_once( + text, + output_emitter=output_emitter, + stream=stream, + timeout=timeout, + turn_index=turn_index, + connection_reused=connection_reused, + reconnected=reconnected, + logical_turn_start=logical_turn_start, + attempts=attempt + 1, + connection_queue_wait_ms=connection_queue_wait_ms, + turn_gap_after_previous_ms=turn_gap_after_previous_ms, + ) + except _XAIEmptyAudioDone as e: + last_error = e + if attempt > 0: + raise APIConnectionError(str(e)) from e + _runtime_logger().warning("XAI_TTS_AUDIO_DONE_WITHOUT_FRAMES | connection_id=%s | segment_id=%s | turn_index=%s | action=retry_same_socket", self._connection_id, stream._segment_id, turn_index) + except _XAIPartialAudioFailure: + raise + except _XAIConnectionClosed as e: + last_error = e + if e.audio_emitted: + confirmed = await self._resynchronize_after_partial_audio(stream=stream, turn_index=turn_index, reason=str(e)) + raise _XAIPartialAudioFailure(f"xAI TTS failed after partial audio; socket_resynchronized={int(confirmed)}: {e}") from e + reconnected = True + await self._reset_ws(reason="before_audio_reconnect") + if attempt > 0: + raise APIConnectionError(str(e)) from e + _runtime_logger().warning( + "XAI_TTS_WS_RECONNECTING | connection_id=%s | segment_id=%s | turn_index=%s | reason=%s", + self._connection_id, + stream._segment_id, + turn_index, + str(e), + ) + except APITimeoutError: + await self._reset_ws() + raise + except APIStatusError: + await self._reset_ws() + raise + + raise APIConnectionError("xAI TTS websocket disconnected") from last_error + + @staticmethod + def _turn_timing( + *, + turn_index: int, + clear_rtt: float, + provider_ttfb: float, + end_to_end_ttfb: float, + max_audio_delta_gap_ms: int, + max_playout_underrun_0ms: int, + connection_reused: bool, + reconnected: bool, + discarded_messages: list[str], + ) -> _TurnTiming: + return _TurnTiming( + turn_index=turn_index, + clear_rtt_ms=round(clear_rtt * 1000, 3), + provider_ttfb_ms=( + round(provider_ttfb * 1000, 3) if provider_ttfb >= 0 else None + ), + end_to_end_ttfb_ms=( + round(end_to_end_ttfb * 1000, 3) if end_to_end_ttfb >= 0 else None + ), + max_audio_delta_gap_ms=max_audio_delta_gap_ms, + max_playout_underrun_0ms=max_playout_underrun_0ms, + connection_reused=connection_reused, + reconnected=reconnected, + discarded_messages=list(discarded_messages), + ) + + async def _resynchronize_after_partial_audio( + self, *, stream: SynthesizeStream, turn_index: int, reason: str + ) -> bool: + clear_sent = False + try: + ws = self._require_ws() + await ws.send_str(json.dumps({"type": "text.clear"})) + self._note_activity() + clear_sent = True + except Exception as exc: + _runtime_logger().warning( + "XAI_TTS_PARTIAL_AUDIO_CLEAR_WRITE_FAILED | connection_id=%s | " + "segment_id=%s | turn_index=%s | reason=%s | error=%s", + self._connection_id, + stream._segment_id, + turn_index, + reason, + type(exc).__name__, + ) + await self._reset_ws( + reason=( + "partial_audio_after_best_effort_clear" + if clear_sent + else "partial_audio_clear_write_failed" + ) + ) + _runtime_logger().warning( + "XAI_TTS_PARTIAL_AUDIO_SOCKET_DISCARDED | connection_id=%s | " + "segment_id=%s | turn_index=%s | reason=%s | clear_sent=%s | " + "action=reconnect_on_next_synthesis", + self._connection_id, + stream._segment_id, + turn_index, + reason, + int(clear_sent), + ) + return False + + async def _synthesize_turn_once( + self, text: str, *, output_emitter: tts.AudioEmitter, stream: SynthesizeStream, + timeout: float, turn_index: int, connection_reused: bool, reconnected: bool, + logical_turn_start: float, attempts: int, + connection_queue_wait_ms: int, + turn_gap_after_previous_ms: int | None, + ) -> _TurnResult: + cache_key = self._owner._take_initial_greeting_capture(text) if self._owner else None + captured_pcm = bytearray() if cache_key is not None else None + ws = self._require_ws() + audio_emitted = clear_acknowledged = False + trace_id: str | None = None + provider_ttfb = end_to_end_ttfb = clear_rtt = -1.0 + discarded_messages: list[str] = [] + provider_audio_bytes = audio_delta_count = empty_audio_deltas = 0 + pcm_duration_s = 0.0 + first_pcm_released_at: float | None = None + last_audio_at: float | None = None + max_gap_ms = max_underrun_ms = micro_underflows = 0 + underflow_episode_max_ms = underflow_episode_total_ms = 0 + underflow_active = False + underflow_error_ms = _underflow_error_ms() + text_delta_sent_at = total_deadline = 0.0 + clear_sequence = self._clear_sequence + 1 + self._clear_sequence = clear_sequence + clear_discarded_message_count = clear_discarded_audio_bytes = 0 + + def observe_underflow_balance(balance_s: float) -> None: + nonlocal max_underrun_ms, micro_underflows, underflow_active + nonlocal underflow_episode_max_ms, underflow_episode_total_ms + debt_ms = max(0, round(-balance_s * 1000)) + max_underrun_ms = max(max_underrun_ms, debt_ms) + if balance_s < 0: + if not underflow_active: + micro_underflows += 1 + underflow_active = True + underflow_episode_max_ms = 0 + underflow_episode_max_ms = max(underflow_episode_max_ms, debt_ms) + elif underflow_active: + underflow_episode_total_ms += underflow_episode_max_ms + underflow_active = False + underflow_episode_max_ms = 0 + + def finish_open_underflow() -> None: + nonlocal underflow_active, underflow_episode_max_ms, underflow_episode_total_ms + if underflow_active: + underflow_episode_total_ms += underflow_episode_max_ms + underflow_active = False + underflow_episode_max_ms = 0 + + def average_underflow_ms() -> int: + return round(underflow_episode_total_ms / micro_underflows) if micro_underflows else 0 + + async def partial_failure(reason: str) -> None: + finish_open_underflow() + provider_synthesis_ms = round( + (time.perf_counter() - text_delta_sent_at) * 1000 + ) + pcm_duration_ms = round(pcm_duration_s * 1000) + average_underrun_ms = average_underflow_ms() + _runtime_logger().warning( + "XAI_TTS_TURN_FAILED | connection_id=%s | segment_id=%s | turn_index=%s | phase=after_playout | reason=%s | attempts=%s | pcm_bytes=%s | pcm_duration_ms=%s | max_audio_delta_gap_ms=%s | xai_underrun_estimado_ms=%s | xai_micro_underflows=%s | xai_avg_underrun_ms=%s | underflow_error_ms=%s", + self._connection_id, stream._segment_id, turn_index, reason, attempts, + provider_audio_bytes, pcm_duration_ms, max_gap_ms, + max_underrun_ms, micro_underflows, average_underrun_ms, underflow_error_ms, + ) + if self._owner is not None: + self._owner.emit( + "xai_tts_turn_failed", + { + "segment_id": stream._segment_id, + "turn_index": turn_index, + "phase": "after_playout", + "reason": reason, + "max_audio_delta_gap_ms": max_gap_ms, + "max_playout_underrun_0ms": max_underrun_ms, + "xai_micro_underflows": micro_underflows, + "xai_avg_underrun_ms": average_underrun_ms, + "provider_synthesis_ms": provider_synthesis_ms, + "provider_ttfb_ms": ( + round(provider_ttfb * 1000, 3) + if provider_ttfb >= 0 + else None + ), + "end_to_end_ttfb_ms": ( + round(end_to_end_ttfb * 1000, 3) + if end_to_end_ttfb >= 0 + else None + ), + "underflow_error_ms": underflow_error_ms, + "pcm_bytes": provider_audio_bytes, + "pcm_duration_ms": pcm_duration_ms, + "attempts": attempts, + "connection_queue_wait_ms": connection_queue_wait_ms, + "turn_gap_after_previous_ms": turn_gap_after_previous_ms, + "connection_reused": connection_reused, + "reconnected": reconnected, + }, + ) + confirmed = await self._resynchronize_after_partial_audio(stream=stream, turn_index=turn_index, reason=reason) + raise _XAIPartialAudioFailure( + "xAI TTS partial audio failure: " + f"{reason}; socket_resynchronized={int(confirmed)}; " + f"xai_underrun_estimado_ms={max_underrun_ms}; " + f"xai_micro_underflows={micro_underflows}; " + f"xai_avg_underrun_ms={average_underrun_ms}; " + f"underflow_error_ms={underflow_error_ms}; " + f"max_audio_delta_gap_ms={max_gap_ms}; " + f"pcm_duration_ms={pcm_duration_ms}; " + f"pcm_bytes={provider_audio_bytes}" + ) + + try: + try: + clear_started = time.perf_counter() + await ws.send_str(json.dumps({"type": "text.clear"})) + self._note_activity() + _runtime_logger().info( + "XAI_TTS_CLEAR_SENT | connection_id=%s | segment_id=%s | turn_index=%s | clear_sequence=%s | connection_reused=%s | queue_wait_ms=%s | gap_after_previous_turn_ms=%s", + self._connection_id, stream._segment_id, turn_index, clear_sequence, + int(connection_reused), connection_queue_wait_ms, turn_gap_after_previous_ms, + ) + text_delta_sent_at = time.perf_counter() + await ws.send_str(json.dumps({"type": "text.delta", "delta": text})) + self._note_activity() + await ws.send_str(json.dumps({"type": "text.done"})) + self._note_activity() + total_timeout_s = _turn_total_timeout_s() + total_deadline = text_delta_sent_at + total_timeout_s + _runtime_logger().info("XAI_TTS_TURN_START | connection_id=%s | segment_id=%s | turn_index=%s | first_frame_timeout_ms=%s | underflow_error_ms=%s | total_timeout_ms=%s | attempts=%s", self._connection_id, stream._segment_id, turn_index, round(_first_frame_timeout_s() * 1000), underflow_error_ms, round(total_timeout_s * 1000), attempts) + except (aiohttp.ClientConnectionError, aiohttp.ClientError, ConnectionResetError, RuntimeError) as exc: + raise _XAIConnectionClosed("xAI TTS websocket clear or text write failed", audio_emitted=False) from exc + stream._mark_started() + while True: + now = time.perf_counter() + total_remaining_s = total_deadline - now + if total_remaining_s <= 0: + if audio_emitted: + observe_underflow_balance( + _estimated_pcm_balance_s( + pcm_duration_s=pcm_duration_s, + first_pcm_released_at=first_pcm_released_at, + now=now, + ) + ) + await partial_failure("total_timeout") + raise _XAIConnectionClosed("xAI TTS total timeout before audio", audio_emitted=False) + if not audio_emitted: + receive_timeout = min(total_remaining_s, _first_frame_timeout_s()) + else: + balance_s = _estimated_pcm_balance_s(pcm_duration_s=pcm_duration_s, first_pcm_released_at=first_pcm_released_at, now=now) + observe_underflow_balance(balance_s) + remaining_to_error_s = balance_s + underflow_error_ms / 1000 + if remaining_to_error_s <= 0: + await partial_failure("underflow_error") + receive_timeout = min(total_remaining_s, remaining_to_error_s) + try: + msg = await asyncio.wait_for(ws.receive(), timeout=receive_timeout) + except asyncio.TimeoutError: + timeout_now = time.perf_counter() + if timeout_now >= total_deadline: + if audio_emitted: + observe_underflow_balance( + _estimated_pcm_balance_s( + pcm_duration_s=pcm_duration_s, + first_pcm_released_at=first_pcm_released_at, + now=timeout_now, + ) + ) + await partial_failure("total_timeout") + raise _XAIConnectionClosed("xAI TTS total timeout before audio", audio_emitted=False) + if not audio_emitted: + raise _XAIConnectionClosed("xAI TTS timed out before audio", audio_emitted=False) + observe_underflow_balance( + _estimated_pcm_balance_s( + pcm_duration_s=pcm_duration_s, + first_pcm_released_at=first_pcm_released_at, + now=timeout_now, + ) + ) + await partial_failure("underflow_error") + self._note_activity() + if msg.type in (aiohttp.WSMsgType.CLOSED, aiohttp.WSMsgType.CLOSE, aiohttp.WSMsgType.CLOSING, aiohttp.WSMsgType.ERROR): + if audio_emitted: + await partial_failure("websocket_closed") + raise _XAIConnectionClosed("xAI TTS websocket closed unexpectedly", audio_emitted=audio_emitted) + if msg.type != aiohttp.WSMsgType.TEXT: + discarded_messages.append(str(msg.type)) + continue + try: + data = json.loads(msg.data) + except json.JSONDecodeError as exc: + if audio_emitted: + await partial_failure("invalid_json") + raise APIStatusError("xAI TTS returned invalid JSON", status_code=-1, request_id=stream._segment_id, body=str(msg.data)) from exc + msg_type = str(data.get("type") or "") + if not clear_acknowledged: + if msg_type == "audio.clear": + clear_acknowledged = True + clear_rtt = time.perf_counter() - clear_started + _runtime_logger().info( + "XAI_TTS_CLEAR_ACK | connection_id=%s | segment_id=%s | turn_index=%s | clear_sequence=%s | clear_rtt_ms=%s | stale_messages=%s | stale_audio_bytes=%s | stale_message_types=%s", + self._connection_id, stream._segment_id, turn_index, clear_sequence, + round(clear_rtt * 1000), clear_discarded_message_count, + clear_discarded_audio_bytes, discarded_messages, + ) + continue + if msg_type == "error": + raise _XAIConnectionClosed("xAI TTS clear failed", audio_emitted=False) + clear_discarded_message_count += 1 + if msg_type == "audio.delta": + try: + clear_discarded_audio_bytes += len( + base64.b64decode(data.get("delta", ""), validate=True) + ) + except (ValueError, TypeError): + pass + discarded_messages.append(msg_type or "unknown") + continue + if msg_type == "audio.delta": + now = time.perf_counter() + try: + payload = base64.b64decode(data.get("delta", ""), validate=True) + except (ValueError, TypeError) as exc: + if audio_emitted: + await partial_failure("invalid_audio_delta") + raise APIStatusError("xAI TTS returned invalid audio.delta", status_code=-1, request_id=stream._segment_id, body=str(data)) from exc + if not payload: + empty_audio_deltas += 1 + _runtime_logger().warning("XAI_TTS_EMPTY_AUDIO_DELTA | connection_id=%s | segment_id=%s | turn_index=%s", self._connection_id, stream._segment_id, turn_index) + continue + if last_audio_at is not None: + max_gap_ms = max(max_gap_ms, round((now - last_audio_at) * 1000)) + if not audio_emitted: + provider_ttfb = now - text_delta_sent_at + end_to_end_ttfb = now - logical_turn_start + first_pcm_released_at = now + stream._note_provider_ttfb(provider_ttfb) + else: + observe_underflow_balance( + _estimated_pcm_balance_s( + pcm_duration_s=pcm_duration_s, + first_pcm_released_at=first_pcm_released_at, + now=now, + ) + ) + provider_audio_bytes += len(payload) + audio_delta_count += 1 + pcm_duration_s += _pcm_duration_s(len(payload)) + observe_underflow_balance( + _estimated_pcm_balance_s( + pcm_duration_s=pcm_duration_s, + first_pcm_released_at=first_pcm_released_at, + now=now, + ) + ) + if captured_pcm is not None: + captured_pcm.extend(payload) + output_emitter.push(payload) + audio_emitted = True + last_audio_at = now + elif msg_type == "audio.done": + finish_open_underflow() + if not audio_emitted: + raise _XAIEmptyAudioDone("xAI TTS audio.done arrived without non-empty PCM") + trace_id = data.get("trace_id") + if cache_key is not None and captured_pcm and self._owner is not None: + cache = self._owner._initial_greeting_audio_cache + if cache is not None: + await cache.store_pcm(cache_key, bytes(captured_pcm)) + timing = replace( + self._turn_timing(turn_index=turn_index, clear_rtt=clear_rtt, provider_ttfb=provider_ttfb, end_to_end_ttfb=end_to_end_ttfb, max_audio_delta_gap_ms=max_gap_ms, max_playout_underrun_0ms=max_underrun_ms, connection_reused=connection_reused, reconnected=reconnected, discarded_messages=discarded_messages), + provider_synthesis_ms=round((time.perf_counter() - text_delta_sent_at) * 1000, 3), + connection_queue_wait_ms=connection_queue_wait_ms, + turn_gap_after_previous_ms=turn_gap_after_previous_ms, + clear_discarded_message_count=clear_discarded_message_count, + clear_discarded_audio_bytes=clear_discarded_audio_bytes, + xai_underrun_estimado_ms=max_underrun_ms, xai_micro_underflows=micro_underflows, xai_avg_underrun_ms=average_underflow_ms(), underflow_error_ms=underflow_error_ms, pcm_bytes=provider_audio_bytes, pcm_duration_ms=round(pcm_duration_s * 1000), empty_audio_deltas=empty_audio_deltas, attempts=attempts, + ) + _runtime_logger().info("XAI_TTS_TURN_DONE | connection_id=%s | segment_id=%s | turn_index=%s | trace_id=%s | provider_ttfb_ms=%s | synthesis_ms=%s | pcm_bytes=%s | pcm_duration_ms=%s | max_audio_delta_gap_ms=%s | xai_underrun_estimado_ms=%s | xai_micro_underflows=%s | xai_avg_underrun_ms=%s | underflow_error_ms=%s | empty_audio_deltas=%s | attempts=%s | clear_rtt_ms=%s | queue_wait_ms=%s | gap_after_previous_turn_ms=%s | stale_messages_before_clear_ack=%s | stale_audio_bytes_before_clear_ack=%s", self._connection_id, stream._segment_id, turn_index, trace_id or "", timing.provider_ttfb_ms, timing.provider_synthesis_ms, timing.pcm_bytes, timing.pcm_duration_ms, timing.max_audio_delta_gap_ms, timing.xai_underrun_estimado_ms, timing.xai_micro_underflows, timing.xai_avg_underrun_ms, timing.underflow_error_ms, timing.empty_audio_deltas, timing.attempts, timing.clear_rtt_ms, timing.connection_queue_wait_ms, timing.turn_gap_after_previous_ms, timing.clear_discarded_message_count, timing.clear_discarded_audio_bytes) + return _TurnResult(audio_emitted=True, trace_id=trace_id, provider_ttfb=provider_ttfb, timing=timing) + elif msg_type == "audio.clear": + if audio_emitted: + await partial_failure("unexpected_audio_clear") + raise _XAIConnectionClosed("xAI TTS received unexpected audio.clear", audio_emitted=False) + elif msg_type == "error": + if audio_emitted: + await partial_failure("provider_error") + raise _XAIConnectionClosed("xAI TTS provider error before audio", audio_emitted=False) + else: + discarded_messages.append(msg_type or "unknown") + except asyncio.CancelledError: + try: + await self._clear_and_wait(timeout=min(timeout, _first_frame_timeout_s())) + except Exception: + await self._reset_ws(reason="cancel_clear_unconfirmed") + raise + + async def _clear_and_wait(self, *, timeout: float) -> tuple[float, list[str]]: + """Clear provider state and wait for its acknowledged clean boundary.""" + ws = self._require_ws() + discarded_messages: list[str] = [] + started_at = time.perf_counter() + try: + await ws.send_str(json.dumps({"type": "text.clear"})) + self._note_activity() + except ( + aiohttp.ClientConnectionError, + aiohttp.ClientError, + ConnectionResetError, + RuntimeError, + ) as e: + await self._reset_ws() + raise _XAIConnectionClosed( + "xAI TTS websocket clear write failed", audio_emitted=False + ) from e + + try: + while True: + msg = await asyncio.wait_for(ws.receive(), timeout=timeout) + self._note_activity() + if msg.type in ( + aiohttp.WSMsgType.CLOSED, + aiohttp.WSMsgType.CLOSE, + aiohttp.WSMsgType.CLOSING, + aiohttp.WSMsgType.ERROR, + ): + await self._reset_ws() + raise _XAIConnectionClosed( + "xAI TTS websocket closed while clearing", audio_emitted=False + ) + if msg.type != aiohttp.WSMsgType.TEXT: + discarded_messages.append(str(msg.type)) + continue + try: + data = json.loads(msg.data) + except json.JSONDecodeError: + discarded_messages.append("invalid-json") + continue + msg_type = str(data.get("type", "unknown")) + if msg_type == "audio.clear": + return time.perf_counter() - started_at, discarded_messages + if msg_type == "error": + raise APIStatusError( + data.get("message", "xAI TTS clear failed"), + status_code=-1, + request_id=data.get("trace_id"), + body=str(data), + ) + discarded_messages.append(msg_type) + except asyncio.TimeoutError: + await self._reset_ws() + raise _XAIConnectionClosed( + "xAI TTS websocket timed out waiting for audio.clear", + audio_emitted=False, + ) from None + + def _require_ws(self) -> aiohttp.ClientWebSocketResponse: + if self._ws is None or self._ws.closed: + raise _XAIConnectionClosed( + "xAI TTS websocket is not connected", + audio_emitted=False, + ) + if self._ws.exception() is not None: + raise _XAIConnectionClosed( + "xAI TTS websocket is in an error state", + audio_emitted=False, + ) + return self._ws + + async def _reset_ws(self, *, reason: str = "reset") -> None: + ws = self._ws + self._ws = None + if ws is not None and not ws.closed: + close_started_at = time.perf_counter() + await ws.close() + _runtime_logger().info( + "XAI_TTS_WS_CLOSED | connection_id=%s | reason=%s | close_ms=%s | age_ms=%s", + self._connection_id, + reason, + _elapsed_ms(close_started_at), + _elapsed_ms(self._opened_at or 0.0), + ) + + async def aclose(self) -> None: + if self._closed: + return + self._closed = True + await self._reset_ws() + + +TTS = OraclexAITTS diff --git a/src/app/livekit/assets/comfort/fails/tts_fail_recovery.wav b/src/app/livekit/assets/comfort/fails/tts_fail_recovery.wav new file mode 100644 index 0000000..1d55cf0 Binary files /dev/null and b/src/app/livekit/assets/comfort/fails/tts_fail_recovery.wav differ diff --git a/src/app/livekit/assets/comfort/interruption/01.txt b/src/app/livekit/assets/comfort/interruption/01.txt new file mode 100644 index 0000000..f59bbe4 --- /dev/null +++ b/src/app/livekit/assets/comfort/interruption/01.txt @@ -0,0 +1 @@ +Um instante diff --git a/src/app/livekit/assets/comfort/interruption/01.wav b/src/app/livekit/assets/comfort/interruption/01.wav new file mode 100644 index 0000000..33f2d3d Binary files /dev/null and b/src/app/livekit/assets/comfort/interruption/01.wav differ diff --git a/src/app/livekit/assets/comfort/long/01.txt b/src/app/livekit/assets/comfort/long/01.txt new file mode 100644 index 0000000..4d4d885 --- /dev/null +++ b/src/app/livekit/assets/comfort/long/01.txt @@ -0,0 +1 @@ +Estou verificando as informações para te ajudar. Só um momentinho. \ No newline at end of file diff --git a/src/app/livekit/assets/comfort/long/01.wav b/src/app/livekit/assets/comfort/long/01.wav new file mode 100644 index 0000000..9c2c687 Binary files /dev/null and b/src/app/livekit/assets/comfort/long/01.wav differ diff --git a/src/app/livekit/assets/comfort/long/02.txt b/src/app/livekit/assets/comfort/long/02.txt new file mode 100644 index 0000000..d226eed --- /dev/null +++ b/src/app/livekit/assets/comfort/long/02.txt @@ -0,0 +1 @@ +Estou confirmando algumas informações. Um instantinho, por favor. \ No newline at end of file diff --git a/src/app/livekit/assets/comfort/long/02.wav b/src/app/livekit/assets/comfort/long/02.wav new file mode 100644 index 0000000..0e917a3 Binary files /dev/null and b/src/app/livekit/assets/comfort/long/02.wav differ diff --git a/src/app/livekit/assets/comfort/long/03.txt b/src/app/livekit/assets/comfort/long/03.txt new file mode 100644 index 0000000..06375ec --- /dev/null +++ b/src/app/livekit/assets/comfort/long/03.txt @@ -0,0 +1 @@ +Peço apenas mais um instante e já retorno com você \ No newline at end of file diff --git a/src/app/livekit/assets/comfort/long/03.wav b/src/app/livekit/assets/comfort/long/03.wav new file mode 100644 index 0000000..9aa4429 Binary files /dev/null and b/src/app/livekit/assets/comfort/long/03.wav differ diff --git a/src/app/livekit/assets/comfort/short/01.txt b/src/app/livekit/assets/comfort/short/01.txt new file mode 100644 index 0000000..988f8c9 --- /dev/null +++ b/src/app/livekit/assets/comfort/short/01.txt @@ -0,0 +1 @@ +Um momentinho \ No newline at end of file diff --git a/src/app/livekit/assets/comfort/short/01.wav b/src/app/livekit/assets/comfort/short/01.wav new file mode 100644 index 0000000..f3b4554 Binary files /dev/null and b/src/app/livekit/assets/comfort/short/01.wav differ diff --git a/src/app/livekit/assets/comfort/short/02.txt b/src/app/livekit/assets/comfort/short/02.txt new file mode 100644 index 0000000..4e5b9aa --- /dev/null +++ b/src/app/livekit/assets/comfort/short/02.txt @@ -0,0 +1 @@ +Um instantinho \ No newline at end of file diff --git a/src/app/livekit/assets/comfort/short/02.wav b/src/app/livekit/assets/comfort/short/02.wav new file mode 100644 index 0000000..d700984 Binary files /dev/null and b/src/app/livekit/assets/comfort/short/02.wav differ diff --git a/src/app/livekit/assets/comfort/short/03.txt b/src/app/livekit/assets/comfort/short/03.txt new file mode 100644 index 0000000..2e48d4f --- /dev/null +++ b/src/app/livekit/assets/comfort/short/03.txt @@ -0,0 +1 @@ +Um momentinho, por favor \ No newline at end of file diff --git a/src/app/livekit/assets/comfort/short/03.wav b/src/app/livekit/assets/comfort/short/03.wav new file mode 100644 index 0000000..7bf84e7 Binary files /dev/null and b/src/app/livekit/assets/comfort/short/03.wav differ diff --git a/src/app/livekit/assets/comfort/short/04.txt b/src/app/livekit/assets/comfort/short/04.txt new file mode 100644 index 0000000..235b933 --- /dev/null +++ b/src/app/livekit/assets/comfort/short/04.txt @@ -0,0 +1 @@ +Um instantinho, por favor. \ No newline at end of file diff --git a/src/app/livekit/assets/comfort/short/04.wav b/src/app/livekit/assets/comfort/short/04.wav new file mode 100644 index 0000000..47b3ec3 Binary files /dev/null and b/src/app/livekit/assets/comfort/short/04.wav differ diff --git a/src/app/livekit/azure_speech.py b/src/app/livekit/azure_speech.py new file mode 100644 index 0000000..f179524 --- /dev/null +++ b/src/app/livekit/azure_speech.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +import os +from typing import Any, Mapping +from urllib.parse import urlsplit, urlunsplit + +AZURE_SPEECH_TTS_SAMPLE_RATE = 16_000 + + +def _pick(*values: Any) -> str: + for value in values: + if value is None: + continue + text = str(value).strip() + if text: + return text + return "" + + +def _normalize_speech_endpoint(value: str) -> str: + endpoint = (value or "").strip() + if not endpoint: + return "" + + parsed = urlsplit(endpoint) + path = (parsed.path or "").strip().rstrip("/") + host = (parsed.hostname or "").strip().lower() + is_custom_domain = host.endswith(".cognitiveservices.azure.com") + + if path in {"", "/"}: + path = "/cognitiveservices/v1" + + if is_custom_domain and path == "/cognitiveservices/v1": + path = "/tts/cognitiveservices/v1" + + return urlunsplit(parsed._replace(path=path)).rstrip("/") + + +def _is_custom_domain_tts_endpoint(value: str) -> bool: + endpoint = (value or "").strip() + if not endpoint: + return False + + parsed = urlsplit(endpoint) + host = (parsed.hostname or "").strip().lower() + path = (parsed.path or "").rstrip("/") + return host.endswith(".cognitiveservices.azure.com") and path in { + "/cognitiveservices/v1", + "/tts/cognitiveservices/v1", + } + + +def resolve_azure_speech_tts_config( + tts_overrides: Mapping[str, Any] | None = None, + *, + environ: Mapping[str, str] | None = None, +) -> tuple[dict[str, str | None], list[str]]: + overrides = dict(tts_overrides or {}) + env = os.environ if environ is None else environ + + speech_key = _pick(env.get("AZURE_SPEECH_KEY")) + speech_auth_token = _pick(env.get("AZURE_SPEECH_AUTH_TOKEN")) + speech_region = _pick(env.get("AZURE_SPEECH_REGION")) + speech_endpoint = _normalize_speech_endpoint( + _pick(env.get("AZURE_SPEECH_ENDPOINT"), env.get("AZURE_SPEECH_HOST")) + ) + voice = _pick(overrides.get("voice"), overrides.get("voice_id"), env.get("AZURE_SPEECH_VOICE")) + language = _pick(overrides.get("language"), env.get("AZURE_SPEECH_LANGUAGE")) + deployment_id = _pick( + overrides.get("deployment_id"), + overrides.get("model_id"), + env.get("AZURE_SPEECH_DEPLOYMENT_ID"), + ) + + missing: list[str] = [] + if not voice: + missing.append("AZURE_SPEECH_VOICE") + if not (speech_endpoint or speech_region): + missing.append("AZURE_SPEECH_ENDPOINT|AZURE_SPEECH_HOST|AZURE_SPEECH_REGION") + if not (speech_key or speech_auth_token): + missing.append("AZURE_SPEECH_KEY|AZURE_SPEECH_AUTH_TOKEN") + + if missing: + return {}, missing + + if speech_endpoint and deployment_id: + parsed = urlsplit(speech_endpoint) + if (parsed.path or "").rstrip("/") == "/tts/cognitiveservices/v1": + speech_endpoint = urlunsplit( + parsed._replace(path="/voice/cognitiveservices/v1") + ).rstrip("/") + + return ( + { + "voice": voice, + "language": language or None, + "speech_key": speech_key or None, + "speech_region": speech_region or None, + "speech_endpoint": speech_endpoint or None, + "deployment_id": deployment_id or None, + "speech_auth_token": speech_auth_token or None, + }, + [], + ) diff --git a/src/app/livekit/call_config.py b/src/app/livekit/call_config.py new file mode 100644 index 0000000..6390a2f --- /dev/null +++ b/src/app/livekit/call_config.py @@ -0,0 +1,10 @@ +from app.common.call_config import ( + normalize_call_config, + resolve_fake_agent_overrides, + resolve_agent_backend_name, + resolve_stt_overrides, + resolve_tts_overrides, + resolve_vad_logging_overrides, + resolve_vad_overrides, + resolve_ws_overrides, +) diff --git a/src/app/livekit/compat.py b/src/app/livekit/compat.py new file mode 100644 index 0000000..63a65c1 --- /dev/null +++ b/src/app/livekit/compat.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +from importlib import import_module + + +def patch_inference_executor_is_alive() -> bool: + """ + Guard LiveKit's health check against a closed multiprocessing.Process. + + livekit-agents==1.3.10 calls InferenceProcExecutor.is_alive() from the + aiohttp health route. After shutdown, multiprocessing raises + ValueError("process object is closed"), which turns a normal unhealthy + state into a 500. + """ + + try: + module = import_module("livekit.agents.ipc.inference_proc_executor") + except ImportError: + return False + + executor_cls = getattr(module, "InferenceProcExecutor", None) + if executor_cls is None: + return False + + current_is_alive = getattr(executor_cls, "is_alive", None) + if current_is_alive is None: + return False + + if getattr(current_is_alive, "__tim_closed_process_guard__", False): + return False + + def _safe_is_alive(self) -> bool: + try: + return current_is_alive(self) + except ValueError as exc: + if str(exc) != "process object is closed": + raise + return False + + _safe_is_alive.__tim_closed_process_guard__ = True + executor_cls.is_alive = _safe_is_alive + return True diff --git a/src/app/livekit/main.py b/src/app/livekit/main.py new file mode 100644 index 0000000..3c7bb59 --- /dev/null +++ b/src/app/livekit/main.py @@ -0,0 +1,1507 @@ +from __future__ import annotations + +import asyncio +import json +import os +import threading +import time +from pathlib import Path +from typing import Any, Callable, Dict, Optional, Tuple + +import httpx +from dotenv import load_dotenv + +from livekit.agents import ( + Agent, + AgentServer, + AgentSession, + JobContext, + JobProcess, + cli, + tts as livekit_tts, + vad as agents_vad, +) +from livekit.plugins import silero, elevenlabs +from livekit.plugins.turn_detector.multilingual import MultilingualModel + +from app.livekit.adapters.agent_backend import AgentBackend, BackendReply +from app.livekit.adapters.azure_rest_tts import AzureRESTTTS +from app.livekit.adapters.backend_factory import build_agent_backend +from app.livekit.runtime.initial_greeting_audio_cache import InitialGreetingAudioCache +from app.livekit.adapters.fake_tts import FakeTTS as FakeSessionTTS +from app.livekit.adapters.xai_tts import OraclexAITTS +from app.livekit.compat import patch_inference_executor_is_alive +from app.livekit.call_config import ( + resolve_agent_backend_name, + resolve_fake_agent_overrides, + resolve_stt_overrides, + resolve_tts_overrides, + resolve_vad_logging_overrides, + resolve_vad_overrides, +) +from app.livekit.azure_speech import ( + AZURE_SPEECH_TTS_SAMPLE_RATE, + resolve_azure_speech_tts_config, +) +from app.livekit.adapters.bridge_gateway import BridgeGateway +from app.livekit.adapters.export_service import ExportService +from app.livekit.adapters.speech_service import ( + SpeechService, + normalize_text, +) +from app.livekit.policies.finalization_policy import FinalizationPolicy +from app.livekit.policies.idle_policy import IdlePolicy, IdlePolicyConfig +from app.livekit.policies.interrupt_policy import InterruptPolicy +from app.livekit.runtime.command_executor import RuntimeCommandExecutor +from app.livekit.runtime.call_runtime import CallRuntime +from app.livekit.runtime.state import RuntimeConfig +from app.livekit.vad_dynamic_threshold import ( + DynamicVADThresholdConfig, + DynamicVADThresholdController, +) +from app.providers.stt_internal_livekit import ( + InternalHTTPSTT, + InternalSTTConfig, + extract_text_from_transcript, +) +from app.providers.stt_fake import FakeSTT +from app.providers.stt_vosk import VoskSTT, DEFAULT_GRAMMAR +from app.utils.call_timeline import CallTimeline +from app.utils.logging import ( + get_call_logger, + log_flow_event, + setup_minimal_logging, + set_log_session_id, + structured_context_from_metadata, +) + +logger = setup_minimal_logging() +load_dotenv() +patch_inference_executor_is_alive() + +# --- ENV / AMBIENTE --- +APP_ENV = (os.getenv("APP_ENV") or os.getenv("ENV") or "dev").strip().lower() +AGENT_BASE_NAME = os.getenv("AGENT_BASE_NAME", "ws-voice-agent").strip() +AGENT_NAME = os.getenv("AGENT_NAME", f"{AGENT_BASE_NAME}-{APP_ENV}").strip() +AGENT_SERVER_PORT = int(os.getenv("AGENT_SERVER_PORT", "18081")) + +server = AgentServer( + num_idle_processes=int(os.getenv("NUM_IDLE_PROCESSES", "5")), + port=AGENT_SERVER_PORT, +) + +CALL_END_GRACE_S = float(os.getenv("CALL_END_GRACE_S", "2")) +EXPORT_SESSION_ID_ENV = os.getenv("EXPORT_SESSION_ID", "").strip() + +# --- TUNING (interrupção nativa) --- +MIN_INTERRUPT_S = float(os.getenv("MIN_INTERRUPT_S", "1.6")) +MIN_ENDPOINTING_DELAY = 1.0 +MAX_ENDPOINTING_DELAY = 1.2 +FINAL_GRACE_S = float(os.getenv("INTERRUPT_FINAL_GRACE_S", "0.12")) +DISCARD_AUDIO_IF_UNINTERRUPTIBLE = ( + os.getenv("DISCARD_AUDIO_IF_UNINTERRUPTIBLE", "1").strip().lower() + in {"1", "true", "yes", "on"} +) +FLOW_LOG_VAD_DECISIONS = ( + os.getenv("FLOW_LOG_VAD_DECISIONS", "0").strip().lower() + in {"1", "true", "yes", "on"} +) +FLOW_LOG_VAD_ACTIVITY = ( + os.getenv("FLOW_LOG_VAD_ACTIVITY", "0").strip().lower() + in {"1", "true", "yes", "on"} +) +FLOW_LOG_VAD_ACTIVITY_MIN_PROB = float(os.getenv("FLOW_LOG_VAD_ACTIVITY_MIN_PROB", "0.03")) +VAD_MIN_SPEECH_DURATION = float(os.getenv("VAD_MIN_SPEECH_DURATION", "0.15")) +VAD_MIN_SILENCE_DURATION = float( + os.getenv("VAD_MIN_SILENCE_DURATION") + or os.getenv("LIVEKIT_VAD_MIN_SILENCE_DURATION_S") + or "0.35" +) +VAD_ACTIVATION_THRESHOLD = float(os.getenv("VAD_ACTIVATION_THRESHOLD", "0.35")) +VAD_DEACTIVATION_THRESHOLD = float( + os.getenv("VAD_DEACTIVATION_THRESHOLD", str(max(VAD_ACTIVATION_THRESHOLD - 0.15, 0.01))) +) +VAD_PREFIX_PADDING_DURATION = float(os.getenv("VAD_PREFIX_PADDING_DURATION", "1.5")) +VAD_PREFIX_PADDING_MIN_DURATION = float(os.getenv("VAD_PREFIX_PADDING_MIN_DURATION", "1.5")) + +DEFAULT_VAD_CONFIG = { + "min_speech_duration": VAD_MIN_SPEECH_DURATION, + "min_silence_duration": VAD_MIN_SILENCE_DURATION, + "activation_threshold": VAD_ACTIVATION_THRESHOLD, + "deactivation_threshold": VAD_DEACTIVATION_THRESHOLD, + "prefix_padding_duration": max(VAD_PREFIX_PADDING_DURATION, VAD_PREFIX_PADDING_MIN_DURATION), +} + +DEFAULT_VAD_LOGGING_CONFIG = { + "log_decisions": FLOW_LOG_VAD_DECISIONS, + "log_activity": FLOW_LOG_VAD_ACTIVITY, + "activity_min_probability": FLOW_LOG_VAD_ACTIVITY_MIN_PROB, +} + +# --- IDLE NUDGE / FECHAMENTO --- +IDLE_NUDGE_ENABLED = (os.getenv("IDLE_NUDGE_ENABLED", "0").strip().lower() in {"1", "true", "yes", "on"}) +IDLE_NUDGE_DELAY_S = float(os.getenv("IDLE_NUDGE_DELAY_S", "30")) +IDLE_NUDGE_JOIN_DELAY_S = float(os.getenv("IDLE_NUDGE_JOIN_DELAY_S", str(IDLE_NUDGE_DELAY_S))) +IDLE_NUDGE_CLOSE_DELAY_S = float(os.getenv("IDLE_NUDGE_CLOSE_DELAY_S", str(IDLE_NUDGE_DELAY_S))) +IDLE_NUDGE_TEXT_ENV = (os.getenv("IDLE_NUDGE_TEXT", "Alô, você está aí?") or "").strip() +IDLE_NUDGE_MAX_TRIES = int(os.getenv("IDLE_NUDGE_MAX_TRIES", "3")) +IDLE_NUDGE_END_REASON = (os.getenv("IDLE_NUDGE_END_REASON", "no_user_response") or "").strip() +AGENT_WAIT_TIMEOUT_RETRY_VAD_THRESHOLD_ENABLED = ( + os.getenv("AGENT_WAIT_TIMEOUT_RETRY_VAD_THRESHOLD_ENABLED", "0").strip().lower() + in {"1", "true", "yes", "on"} +) +AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD = float( + os.getenv("AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD", "0.20") +) +_AGENT_WAIT_TIMEOUT_RETRY_VAD_DEACTIVATION_THRESHOLD_RAW = ( + os.getenv("AGENT_WAIT_TIMEOUT_RETRY_VAD_DEACTIVATION_THRESHOLD", "") or "" +).strip() +AGENT_WAIT_TIMEOUT_RETRY_VAD_DEACTIVATION_THRESHOLD = ( + float(_AGENT_WAIT_TIMEOUT_RETRY_VAD_DEACTIVATION_THRESHOLD_RAW) + if _AGENT_WAIT_TIMEOUT_RETRY_VAD_DEACTIVATION_THRESHOLD_RAW + else None +) +REMOTE_AGENT_INFLIGHT_WAIT_INTERVAL_S = float(os.getenv("REMOTE_AGENT_INFLIGHT_WAIT_INTERVAL_S", "12")) +REMOTE_AGENT_INFLIGHT_WAIT_TIMEOUT_S = float(os.getenv("REMOTE_AGENT_INFLIGHT_WAIT_TIMEOUT_S", "180")) +# Zero means unlimited notices. The total wait is bounded exclusively by the +# backend timeout above, preventing a silent gap between a notice-count limit +# and termination. +REMOTE_AGENT_INFLIGHT_WAIT_MAX_NOTICES = int(os.getenv("REMOTE_AGENT_INFLIGHT_WAIT_MAX_NOTICES", "0")) +REMOTE_AGENT_INFLIGHT_WAIT_TEXT = ( + os.getenv("REMOTE_AGENT_INFLIGHT_WAIT_TEXT", "Um momento, ainda estou consultando para te ajudar.") + or "" +).strip() +PRE_BACKEND_WAIT_NOTICE_FAST_ON_VAD_PAUSE = ( + os.getenv("PRE_BACKEND_WAIT_NOTICE_FAST_ON_VAD_PAUSE", "0").strip().lower() + in {"1", "true", "yes", "on"} +) +LIVEKIT_APP_DIR = Path(__file__).resolve().parent +DEFERRED_INTERRUPTION_MIN_AUDIO_MS = max( + 0, + int(os.getenv("DEFERRED_INTERRUPTION_MIN_AUDIO_MS", "1000")), +) +DEFERRED_INTERRUPTION_ENABLED = ( + os.getenv("DEFERRED_INTERRUPTION_ENABLED", "1").strip().lower() + in {"1", "true", "yes", "on"} +) +DEFERRED_INTERRUPTION_STT_SETTLE_TIMEOUT_S = max( + 0.0, + float(os.getenv("DEFERRED_INTERRUPTION_STT_SETTLE_TIMEOUT_S", "3.0")), +) +DEFERRED_INTERRUPTION_USER_TURN_TIMEOUT_S = max( + 0.0, + float(os.getenv("DEFERRED_INTERRUPTION_USER_TURN_TIMEOUT_S", "10.0")), +) +REMOTE_AGENT_INFLIGHT_WAIT_AUDIO_BASE_DIR = LIVEKIT_APP_DIR / "assets" / "comfort" +REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR = ( + os.getenv( + "REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR", + str(REMOTE_AGENT_INFLIGHT_WAIT_AUDIO_BASE_DIR / "short"), + ) + or "" +).strip() +REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR = ( + os.getenv( + "REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR", + str(REMOTE_AGENT_INFLIGHT_WAIT_AUDIO_BASE_DIR / "long"), + ) + or "" +).strip() +REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_PATH = ( + os.getenv( + "REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_PATH", + "", + ) + or "" +).strip() +REMOTE_AGENT_INFLIGHT_WAIT_AUDIO_ENABLED = bool( + REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR + or REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR + or REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_PATH +) + + +def _float_override( + value: Any, + default: float, + *, + min_value: Optional[float] = None, + max_value: Optional[float] = None, +) -> float: + if value in (None, ""): + return default + try: + resolved = float(value) + except (TypeError, ValueError): + return default + if min_value is not None and resolved < min_value: + return default + if max_value is not None and resolved > max_value: + return default + return resolved + + +def _bool_override(value: Any, default: bool) -> bool: + if value in (None, ""): + return default + raw = str(value).strip().lower() + if raw in {"1", "true", "yes", "on"}: + return True + if raw in {"0", "false", "no", "off"}: + return False + return default + + +def _has_any_override(overrides: Dict[str, Any]) -> bool: + return any(value not in (None, "") for value in overrides.values()) + + +def resolve_vad_runtime_config(call_config: Optional[Dict[str, Any]]) -> Tuple[Dict[str, float], bool]: + overrides = resolve_vad_overrides(call_config) + vad_tuning_overrides = { + key: overrides.get(key) + for key in ( + "min_speech_duration", + "min_silence_duration", + "activation_threshold", + "deactivation_threshold", + "prefix_padding_duration", + ) + } + return ( + { + "min_speech_duration": _float_override( + overrides.get("min_speech_duration"), + DEFAULT_VAD_CONFIG["min_speech_duration"], + min_value=0.0, + ), + "min_silence_duration": _float_override( + overrides.get("min_silence_duration"), + DEFAULT_VAD_CONFIG["min_silence_duration"], + min_value=0.0, + ), + "activation_threshold": _float_override( + overrides.get("activation_threshold"), + DEFAULT_VAD_CONFIG["activation_threshold"], + min_value=0.0, + max_value=1.0, + ), + "deactivation_threshold": _float_override( + overrides.get("deactivation_threshold"), + DEFAULT_VAD_CONFIG["deactivation_threshold"], + min_value=0.0, + max_value=1.0, + ), + "prefix_padding_duration": _float_override( + overrides.get("prefix_padding_duration"), + DEFAULT_VAD_CONFIG["prefix_padding_duration"], + min_value=VAD_PREFIX_PADDING_MIN_DURATION, + ), + }, + _has_any_override(vad_tuning_overrides), + ) + + +def resolve_vad_logging_runtime_config(call_config: Optional[Dict[str, Any]]) -> Dict[str, Any]: + overrides = resolve_vad_logging_overrides(call_config) + return { + "log_decisions": _bool_override( + overrides.get("log_decisions"), + DEFAULT_VAD_LOGGING_CONFIG["log_decisions"], + ), + "log_activity": _bool_override( + overrides.get("log_activity"), + DEFAULT_VAD_LOGGING_CONFIG["log_activity"], + ), + "activity_min_probability": _float_override( + overrides.get("activity_min_probability"), + DEFAULT_VAD_LOGGING_CONFIG["activity_min_probability"], + min_value=0.0, + max_value=1.0, + ), + } + + +def resolve_agent_runtime_config(call_config: Optional[Dict[str, Any]]) -> Tuple[Dict[str, Any], bool]: + overrides = resolve_vad_overrides(call_config) + fast_on_vad_pause_override = overrides.get("pre_backend_wait_notice_fast_on_vad_pause") + deferred_interruption_min_audio_ms_override = overrides.get( + "deferred_interruption_min_audio_ms" + ) + deferred_interruption_enabled_override = overrides.get( + "deferred_interruption_enabled" + ) + return ( + { + "pre_backend_wait_notice_fast_on_vad_pause": _bool_override( + fast_on_vad_pause_override, + PRE_BACKEND_WAIT_NOTICE_FAST_ON_VAD_PAUSE, + ), + "deferred_interruption_min_audio_ms": int( + _float_override( + deferred_interruption_min_audio_ms_override, + float(DEFERRED_INTERRUPTION_MIN_AUDIO_MS), + min_value=0.0, + ) + ), + "deferred_interruption_enabled": _bool_override( + deferred_interruption_enabled_override, + DEFERRED_INTERRUPTION_ENABLED, + ), + }, + any( + value not in (None, "") + for value in ( + fast_on_vad_pause_override, + deferred_interruption_min_audio_ms_override, + deferred_interruption_enabled_override, + ) + ), + ) + + +def load_silero_vad(vad_config: Dict[str, float]) -> agents_vad.VAD: + return silero.VAD.load( + min_speech_duration=vad_config["min_speech_duration"], + min_silence_duration=vad_config["min_silence_duration"], + activation_threshold=vad_config["activation_threshold"], + deactivation_threshold=vad_config["deactivation_threshold"], + prefix_padding_duration=vad_config["prefix_padding_duration"], + ) + + +def vad_config_cache_key(vad_config: Dict[str, float]) -> Tuple[float, float, float, float, float]: + return ( + vad_config["min_speech_duration"], + vad_config["min_silence_duration"], + vad_config["activation_threshold"], + vad_config["deactivation_threshold"], + vad_config["prefix_padding_duration"], + ) + + +def require_env(name: str, *, context: str, timeline: Optional[CallTimeline] = None) -> str: + value = (os.getenv(name, "") or "").strip() + if value: + return value + + if timeline is not None: + timeline.emit( + "config_missing", + key=name, + context=context, + ) + + raise RuntimeError( + f"Missing required env var {name} for {context}. " + f"Add {name} to .env.dev before starting the agent." + ) + + +def create_task_logged(coro, *, name: str): + task = asyncio.create_task(coro, name=name) + + def _done(t: asyncio.Task): + try: + exc = t.exception() + except asyncio.CancelledError: + return + except Exception: + logger.exception("[task:%s] erro lendo exception()", t.get_name()) + return + if exc: + logger.exception("[task:%s] exception: %r", t.get_name(), exc) + + task.add_done_callback(_done) + return task + + +class FlowLoggingVAD(agents_vad.VAD): + def __init__( + self, + inner: agents_vad.VAD, + *, + call_logger: Any, + min_interrupt_s: float, + vad_config: Dict[str, float], + log_decisions: bool, + log_activity: bool, + activity_min_prob: float, + threshold_provider: Optional[Callable[[], Dict[str, Any]]] = None, + on_pre_backend_wait_pause: Optional[Callable[[str], None]] = None, + on_speech_end: Optional[Callable[[int], None]] = None, + ) -> None: + super().__init__(capabilities=inner.capabilities) + self._inner = inner + self._call_logger = call_logger + self._min_interrupt_s = min_interrupt_s + self._vad_config = dict(vad_config) + self._log_decisions = log_decisions + self._log_activity = log_activity + self._activity_min_prob = activity_min_prob + self._threshold_provider = threshold_provider + self._on_pre_backend_wait_pause = on_pre_backend_wait_pause + self._on_speech_end = on_speech_end + self._stream_lock = threading.Lock() + self._next_stream_id = 0 + self._logging_stream_id: Optional[int] = None + + def _claim_logging_stream(self) -> Tuple[int, bool]: + with self._stream_lock: + stream_id = self._next_stream_id + self._next_stream_id += 1 + if self._logging_stream_id is None: + self._logging_stream_id = stream_id + return stream_id, True + return stream_id, False + + def _release_logging_stream(self, stream_id: int) -> None: + with self._stream_lock: + if self._logging_stream_id == stream_id: + self._logging_stream_id = None + + @property + def model(self) -> str: + return self._inner.model + + @property + def provider(self) -> str: + return self._inner.provider + + def stream(self): + stream_id, should_log = self._claim_logging_stream() + return FlowLoggingVADStream( + self._inner.stream(), + stream_id=stream_id, + should_log=should_log, + release_logging_stream=self._release_logging_stream, + call_logger=self._call_logger, + min_interrupt_s=self._min_interrupt_s, + vad_config=self._vad_config, + log_decisions=self._log_decisions, + log_activity=self._log_activity, + activity_min_prob=self._activity_min_prob, + threshold_provider=self._threshold_provider, + on_pre_backend_wait_pause=self._on_pre_backend_wait_pause, + on_speech_end=self._on_speech_end, + ) + + +class FlowLoggingVADStream: + def __init__( + self, + inner: Any, + *, + stream_id: int, + should_log: bool, + release_logging_stream: Callable[[int], None], + call_logger: Any, + min_interrupt_s: float, + vad_config: Dict[str, float], + log_decisions: bool, + log_activity: bool, + activity_min_prob: float, + threshold_provider: Optional[Callable[[], Dict[str, Any]]] = None, + on_pre_backend_wait_pause: Optional[Callable[[str], None]] = None, + on_speech_end: Optional[Callable[[int], None]] = None, + ) -> None: + self._inner = inner + self._stream_id = stream_id + self._should_log = should_log + self._release_logging_stream = release_logging_stream + self._call_logger = call_logger + self._min_interrupt_s = min_interrupt_s + self._vad_config = dict(vad_config) + self._log_decisions = log_decisions + self._log_activity = log_activity + self._activity_min_prob = activity_min_prob + self._threshold_provider = threshold_provider + self._on_pre_backend_wait_pause = on_pre_backend_wait_pause + self._on_speech_end = on_speech_end + self._last_activity_log_at = 0.0 + self._reset_speech_log_state() + + def push_frame(self, frame): + return self._inner.push_frame(frame) + + def flush(self) -> None: + return self._inner.flush() + + def end_input(self) -> None: + return self._inner.end_input() + + async def aclose(self) -> None: + try: + return await self._inner.aclose() + finally: + self._release_logging_stream(self._stream_id) + + def __aiter__(self): + return self + + async def __anext__(self): + ev = await self._inner.__anext__() + self._log_vad_decision(ev) + return ev + + def _reset_speech_log_state(self) -> None: + self._threshold_logged = False + self._activity_logged = False + self._pause_logged = False + self._pre_backend_wait_pause_notified = False + self._speech_seen = False + self._max_speech_ms = 0 + self._max_silence_ms = 0 + + @staticmethod + def _duration_ms(value: Any) -> int: + return round(float(value or 0.0) * 1000) + + def _pre_backend_wait_min_speech_ms(self) -> int: + return max(1, round(float(self._vad_config["min_speech_duration"]) * 1000)) + + def _current_thresholds(self) -> Dict[str, Any]: + provider = self._threshold_provider + if provider is not None: + try: + thresholds = provider() + except Exception: + logger.exception("[vad] failed reading dynamic threshold state") + else: + if isinstance(thresholds, dict): + return { + "mode": thresholds.get("mode") or "dynamic", + "activation_threshold": float( + thresholds.get("activation_threshold", self._vad_config["activation_threshold"]) + ), + "deactivation_threshold": float( + thresholds.get( + "deactivation_threshold", + self._vad_config["deactivation_threshold"], + ) + ), + } + return { + "mode": "baseline", + "activation_threshold": self._vad_config["activation_threshold"], + "deactivation_threshold": self._vad_config["deactivation_threshold"], + } + + def _observe_speech_ms(self, *, speech_ms: int, raw_speech_ms: int) -> int: + self._max_speech_ms = max(self._max_speech_ms, speech_ms, raw_speech_ms) + return self._max_speech_ms + + def _mark_speech_seen_if_meaningful( + self, + *, + speaking: bool, + speech_ms: int, + raw_speech_ms: int, + ) -> None: + observed_speech_ms = self._observe_speech_ms( + speech_ms=speech_ms, + raw_speech_ms=raw_speech_ms, + ) + if observed_speech_ms >= self._pre_backend_wait_min_speech_ms(): + self._speech_seen = True + + def _observed_silence_ms(self, *, raw_silence_ms: int, silence_ms: int) -> int: + self._max_silence_ms = max(self._max_silence_ms, raw_silence_ms, silence_ms) + return self._max_silence_ms + + def _log_user_pause( + self, + *, + decision: str, + silence_ms: int, + pause_min_ms: int, + speech_ms: int, + min_ms: int, + eligible: bool, + probability: float, + raw_speech_ms: int, + raw_silence_ms: int, + ) -> None: + self._pause_logged = True + log_flow_event( + self._call_logger, + "vad_user_pause", + decision=decision, + silence_ms=silence_ms, + pause_min_ms=pause_min_ms, + speech_duration_ms=speech_ms, + min_interrupt_ms=min_ms, + eligible=eligible, + probability=f"{probability:.3f}", + raw_speech_ms=raw_speech_ms, + raw_silence_ms=raw_silence_ms, + ) + + def _maybe_notify_pre_backend_wait_end_of_speech( + self, + *, + silence_ms: int, + speech_ms: int, + pause_min_ms: int, + ) -> None: + callback = self._on_pre_backend_wait_pause + if ( + callback is None + or self._pre_backend_wait_pause_notified + or not self._speech_seen + or self._max_speech_ms < self._pre_backend_wait_min_speech_ms() + ): + return + + self._pre_backend_wait_pause_notified = True + log_flow_event( + self._call_logger, + "pre_backend_wait_end_of_speech", + silence_ms=silence_ms, + pause_min_ms=pause_min_ms, + speech_duration_ms=speech_ms, + ) + try: + callback("vad_end_of_speech") + except Exception: + logger.exception("[vad] failed notifying pre backend wait pause") + + def _log_vad_decision(self, ev: agents_vad.VADEvent) -> None: + if not self._should_log: + return + + min_ms = round(self._min_interrupt_s * 1000) + speech_ms = self._duration_ms(getattr(ev, "speech_duration", 0.0)) + eligible = speech_ms >= min_ms + probability = float(getattr(ev, "probability", 0.0) or 0.0) + raw_speech_ms = self._duration_ms(getattr(ev, "raw_accumulated_speech", 0.0)) + silence_ms = self._duration_ms(getattr(ev, "silence_duration", 0.0)) + raw_silence_ms = self._duration_ms(getattr(ev, "raw_accumulated_silence", 0.0)) + current_silence_ms = max(raw_silence_ms, silence_ms) + pause_min_ms = round(float(self._vad_config["min_silence_duration"]) * 1000) + + if ev.type == agents_vad.VADEventType.START_OF_SPEECH: + self._reset_speech_log_state() + self._mark_speech_seen_if_meaningful( + speaking=True, + speech_ms=speech_ms, + raw_speech_ms=raw_speech_ms, + ) + if self._log_decisions: + log_flow_event( + self._call_logger, + "vad_speech_start", + speech_duration_ms=speech_ms, + min_interrupt_ms=min_ms, + ) + return + + if ev.type == agents_vad.VADEventType.INFERENCE_DONE: + speaking = bool(getattr(ev, "speaking", False)) + self._mark_speech_seen_if_meaningful( + speaking=speaking, + speech_ms=speech_ms, + raw_speech_ms=raw_speech_ms, + ) + if current_silence_ms < 200: + self._pause_logged = False + self._max_silence_ms = 0 + if not self._speech_seen: + self._max_speech_ms = 0 + observed_silence_ms = self._observed_silence_ms( + raw_silence_ms=raw_silence_ms, + silence_ms=silence_ms, + ) + if self._log_activity: + activity_detected = speaking or raw_speech_ms > 0 or probability >= self._activity_min_prob + should_log_activity = activity_detected and ( + not self._activity_logged + or eligible + or (time.monotonic() - self._last_activity_log_at) >= 0.25 + ) + if should_log_activity: + thresholds = self._current_thresholds() + self._activity_logged = True + self._last_activity_log_at = time.monotonic() + log_flow_event( + self._call_logger, + "vad_activity", + speaking=speaking, + speech_duration_ms=speech_ms, + min_interrupt_ms=min_ms, + eligible=eligible, + probability=f"{probability:.3f}", + threshold_mode=thresholds["mode"], + activation_threshold=f"{thresholds['activation_threshold']:.3f}", + deactivation_threshold=f"{thresholds['deactivation_threshold']:.3f}", + activity_log_min_prob=f"{self._activity_min_prob:.3f}", + silence_ms=observed_silence_ms, + pause_min_ms=pause_min_ms, + silence_remaining_ms=max(0, pause_min_ms - observed_silence_ms), + raw_speech_ms=raw_speech_ms, + raw_silence_ms=raw_silence_ms, + ) + elif not activity_detected and raw_silence_ms >= 200: + self._activity_logged = False + + if ( + self._speech_seen + and observed_silence_ms >= pause_min_ms + ): + if not self._pause_logged: + self._log_user_pause( + decision="pause_threshold_reached", + silence_ms=observed_silence_ms, + pause_min_ms=pause_min_ms, + speech_ms=speech_ms, + min_ms=min_ms, + eligible=eligible, + probability=probability, + raw_speech_ms=raw_speech_ms, + raw_silence_ms=raw_silence_ms, + ) + self._maybe_notify_pre_backend_wait_end_of_speech( + silence_ms=observed_silence_ms, + speech_ms=speech_ms, + pause_min_ms=pause_min_ms, + ) + + if self._log_decisions and speaking and eligible and not self._threshold_logged: + self._threshold_logged = True + log_flow_event( + self._call_logger, + "vad_interrupt_check", + decision="eligible_by_duration", + speech_duration_ms=speech_ms, + min_interrupt_ms=min_ms, + probability=f"{probability:.3f}", + raw_speech_ms=raw_speech_ms, + raw_silence_ms=raw_silence_ms, + ) + return + + if ev.type == agents_vad.VADEventType.END_OF_SPEECH: + self._observe_speech_ms( + speech_ms=speech_ms, + raw_speech_ms=raw_speech_ms, + ) + if self._on_speech_end is not None: + try: + # Silero's final speech_duration excludes the terminal + # silence used to detect the endpoint. Do not use the + # maximum INFERENCE_DONE duration (which still includes + # that silence) or ev.frames (which include VAD padding). + self._on_speech_end(speech_ms) + except Exception: + logger.exception("[vad] failed notifying speech duration") + observed_silence_ms = self._observed_silence_ms( + raw_silence_ms=raw_silence_ms, + silence_ms=silence_ms, + ) + self._maybe_notify_pre_backend_wait_end_of_speech( + silence_ms=observed_silence_ms, + speech_ms=speech_ms, + pause_min_ms=pause_min_ms, + ) + if self._speech_seen and not self._pause_logged: + self._log_user_pause( + decision="end_of_speech", + silence_ms=observed_silence_ms, + pause_min_ms=pause_min_ms, + speech_ms=speech_ms, + min_ms=min_ms, + eligible=eligible, + probability=probability, + raw_speech_ms=raw_speech_ms, + raw_silence_ms=raw_silence_ms, + ) + if self._log_decisions: + log_flow_event( + self._call_logger, + "vad_speech_end", + decision="eligible_by_duration" if eligible else "too_short", + speech_duration_ms=speech_ms, + min_interrupt_ms=min_ms, + eligible=eligible, + silence_ms=observed_silence_ms, + raw_silence_ms=raw_silence_ms, + pause_min_ms=pause_min_ms, + ) + self._reset_speech_log_state() + + +class PreBackendWaitNoticeBridge: + def __init__(self) -> None: + self._runtime: Optional[CallRuntime] = None + + def bind(self, runtime: CallRuntime) -> None: + self._runtime = runtime + + def notify(self, reason: str) -> None: + runtime = self._runtime + if runtime is None: + return + runtime.schedule_pre_backend_wait_notice(reason=reason) + + def notify_speech_end(self, speech_duration_ms: int) -> None: + runtime = self._runtime + if runtime is None: + return + runtime.note_vad_speech_end(speech_duration_ms) + + +def prewarm(proc: JobProcess): + proc.userdata["vad"] = load_silero_vad(DEFAULT_VAD_CONFIG) + proc.userdata["vad_cache"] = { + vad_config_cache_key(DEFAULT_VAD_CONFIG): proc.userdata["vad"], + } + proc.userdata["initial_greeting_audio_cache"] = InitialGreetingAudioCache() + wait_audio_loaded, wait_audio_failed = CallRuntime.prewarm_wait_audio_cache( + short_audio_dir=REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR, + long_audio_dir=REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR, + long_audio_path=REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_PATH, + logger=logger, + ) + if wait_audio_loaded or wait_audio_failed: + logger.info( + "WAIT_AUDIO_PREWARM | loaded=%s | failed=%s", + wait_audio_loaded, + wait_audio_failed, + ) + + proc.userdata["http"] = httpx.AsyncClient( + limits=httpx.Limits(max_connections=50, max_keepalive_connections=20), + timeout=httpx.Timeout(30.0), + ) + + vosk_model_path = os.getenv("VOSK_MODEL_PATH", "").strip() + if vosk_model_path: + use_grammar = os.getenv("VOSK_USE_GRAMMAR", "1") == "1" + grammar = DEFAULT_GRAMMAR if use_grammar else None + + proc.userdata["vosk"] = VoskSTT( + model_path=vosk_model_path, + sample_rate=int(os.getenv("VOSK_SAMPLE_RATE", "16000")), + grammar_json=grammar, + ) + proc.userdata["vosk_lock"] = threading.Lock() + else: + proc.userdata["vosk"] = None + proc.userdata["vosk_lock"] = None + logger.debug("[prewarm] VOSK_MODEL_PATH vazio; vosk desativado") + + +server.setup_fnc = prewarm + +class PipelineVoiceAgent(Agent): + def __init__( + self, + *, + elegibility: bool, + protocol: str, + intro: str, + backend_name: str = "", + remote_agent_context: Optional[Dict[str, Any]] = None, + timeline: Optional[CallTimeline] = None, + ) -> None: + super().__init__(instructions=( + "Você é um agente de voz em português (pt-BR). " + "Responda curto, direto e educado. " + "Não use emojis, markdown ou caracteres especiais." + )) + self._elegibility = elegibility + self._protocol = protocol + self._intro = intro + self._backend_name = backend_name + self._remote_agent_context = remote_agent_context or {} + self._timeline = timeline + + self.pipeline: Optional[AgentBackend] = None + self._ready = asyncio.Event() + self._prepare_task: Optional[asyncio.Task] = None + self._run_lock = asyncio.Lock() + + self._end_lock = asyncio.Lock() + self._ended = False + self._end_reply: Optional[BackendReply] = None + + self._pending_interrupt: Optional[Tuple[str, str, bool]] = None + self._pending_lock = asyncio.Lock() + + async def set_pending_interrupt( + self, + *, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + listened_text = normalize_text(listened_text or "") + speech_id = str(speech_id or "").strip() + async with self._pending_lock: + self._pending_interrupt = (speech_id, listened_text, bool(skipped)) + + async def consume_pending_interrupt(self) -> Tuple[Optional[str], str, bool]: + async with self._pending_lock: + v = self._pending_interrupt + self._pending_interrupt = None + if not v: + return None, "", False + speech_id, listened_text, skipped = v + return (normalize_text(listened_text or ""), speech_id, bool(skipped)) + + async def on_enter(self): + if self._timeline is not None: + self._timeline.emit( + "agent_backend_build_started", + backend=self._backend_name or "remote_ws", + ) + self.pipeline = build_agent_backend( + intro=self._intro, + backend_name=self._backend_name, + remote_agent_context=self._remote_agent_context, + timeline=self._timeline, + streaming=False, + ) + + self._prepare_task = asyncio.create_task( + self.pipeline.prepare(self._elegibility, self._protocol) + ) + + prepare_ok = False + try: + await self._prepare_task + prepare_ok = True + except Exception: + logger.exception("[pipeline] prepare() falhou") + if self._timeline is not None: + self._timeline.emit( + "agent_backend_prepare_failed", + backend=self._backend_name or "remote_ws", + ) + finally: + if self._timeline is not None and prepare_ok: + self._timeline.emit( + "agent_backend_ready", + backend=self._backend_name or "remote_ws", + ) + self._ready.set() + + async def end_service_once(self) -> BackendReply: + async with self._end_lock: + if self._ended: + return self._end_reply or BackendReply(stage="DONE", done=True, export_payload=[]) + + self._ended = True + if self.pipeline is None: + self._end_reply = BackendReply(stage="DONE", done=True, export_payload=[]) + return self._end_reply + + try: + self._end_reply = await self.pipeline.end_service_once() + except Exception: + logger.exception("[pipeline] end_service() falhou") + self._end_reply = BackendReply(stage="DONE", done=True, export_payload=[]) + + return self._end_reply + + +@server.rtc_session(agent_name=AGENT_NAME) +async def entrypoint(ctx: JobContext): + ctx.log_context_fields = {"room": ctx.room.name} + + http_client: httpx.AsyncClient = ctx.proc.userdata["http"] + default_vad = ctx.proc.userdata["vad"] + vad = default_vad + vosk = ctx.proc.userdata.get("vosk") + initial_greeting_audio_cache = ctx.proc.userdata["initial_greeting_audio_cache"] + vosk_lock = ctx.proc.userdata.get("vosk_lock") + + try: + md = json.loads(ctx.job.metadata or "{}") + except Exception: + md = {} + + session_data = md.get("session_data") or {} + elegibility = bool(md.get("elegibility", True)) + protocol = str(md.get("protocol", os.getenv("DEFAULT_PROTOCOL", "PRT-2025-0000123456"))) + intro = md.get("intro", "") + agent_starts_conversation = bool(md.get("agent_starts_conversation", False)) + bridge_identity = str(md.get("bridge_identity", "")).strip() + call_config = md.get("call_config") or {} + remote_agent_context = dict(md.get("remote_agent") or {}) + fake_agent_overrides = resolve_fake_agent_overrides(call_config) + if fake_agent_overrides.get("delay_ms") is not None: + remote_agent_context["_fake_agent_delay_ms"] = fake_agent_overrides["delay_ms"] + if fake_agent_overrides.get("responses"): + remote_agent_context["_fake_agent_responses"] = fake_agent_overrides["responses"] + debug_events_enabled = bool(md.get("debug_events_enabled", False)) + timeline_origin_ms = int(md.get("timeline_origin_ms") or round(time.time() * 1000)) + structured_context = structured_context_from_metadata(md) + session_id = str( + md.get("session_id") + or session_data.get("session_id") + or session_data.get("sessionId") + or remote_agent_context.get("session_id") + or remote_agent_context.get("sessionId") + or "" + ).strip() + set_log_session_id(session_id) + ctx.log_context_fields = {"room": ctx.room.name, "session_id": session_id} + + phone_number = ( + session_data.get("gsm") + or session_data.get("msisdn") + or session_data.get("phone") + or remote_agent_context.get("GSM") + or "" + ) + + call_logger = get_call_logger( + phone_number=phone_number, + session_id=session_id, + ) + timeline = CallTimeline( + logger=call_logger, + component="agent", + timeline_id=str(md.get("timeline_id") or ctx.room.name), + protocol=protocol, + room=ctx.room.name, + session_id=session_id, + phone_number=phone_number, + origin_unix_ms=timeline_origin_ms, + ) + + vad_config, has_vad_overrides = resolve_vad_runtime_config(call_config) + vad_logging_config = resolve_vad_logging_runtime_config(call_config) + agent_runtime_config, has_agent_runtime_overrides = resolve_agent_runtime_config(call_config) + vad_config_source = "call_config" if has_vad_overrides else "env" + pre_backend_wait_notice_fast_source = ( + "call_config" if has_agent_runtime_overrides else "env" + ) + if has_vad_overrides: + try: + vad_cache = ctx.proc.userdata.setdefault("vad_cache", {}) + cache_key = vad_config_cache_key(vad_config) + vad = vad_cache.get(cache_key) + if vad is None: + vad = load_silero_vad(vad_config) + vad_cache[cache_key] = vad + timeline.emit( + "vad_config_override_applied", + min_speech_duration=vad_config["min_speech_duration"], + min_silence_duration=vad_config["min_silence_duration"], + activation_threshold=vad_config["activation_threshold"], + deactivation_threshold=vad_config["deactivation_threshold"], + prefix_padding_duration=vad_config["prefix_padding_duration"], + ) + except Exception: + logger.exception("[vad] failed loading call_config override; falling back to env VAD") + call_logger.exception("[vad] failed loading call_config override; falling back to env VAD") + vad = default_vad + vad_config = dict(DEFAULT_VAD_CONFIG) + vad_config_source = "env_fallback" + timeline.emit("vad_config_override_failed") + + agent_wait_timeout_retry_vad_threshold_enabled = ( + AGENT_WAIT_TIMEOUT_RETRY_VAD_THRESHOLD_ENABLED + and AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD + < vad_config["activation_threshold"] + ) + agent_wait_timeout_retry_vad_deactivation_threshold = min( + AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD, + max( + 0.01, + ( + AGENT_WAIT_TIMEOUT_RETRY_VAD_DEACTIVATION_THRESHOLD + if AGENT_WAIT_TIMEOUT_RETRY_VAD_DEACTIVATION_THRESHOLD is not None + else vad_config["deactivation_threshold"] + ), + ), + ) + + if agent_wait_timeout_retry_vad_threshold_enabled: + try: + vad = load_silero_vad(vad_config) + vad_config_source = f"{vad_config_source}_dedicated" + timeline.emit( + "vad_dynamic_threshold_dedicated_instance", + activation_threshold=vad_config["activation_threshold"], + deactivation_threshold=vad_config["deactivation_threshold"], + ) + except Exception: + logger.exception("[vad] failed loading dedicated VAD for dynamic threshold") + call_logger.exception("[vad] failed loading dedicated VAD for dynamic threshold") + agent_wait_timeout_retry_vad_threshold_enabled = False + timeline.emit("vad_dynamic_threshold_dedicated_instance_failed") + + nudge_text = (str(md.get("nudge") or IDLE_NUDGE_TEXT_ENV or "Alô, você está aí?") or "").strip() + + call_t0 = time.monotonic() + + await ctx.connect() + call_logger.info( + "CALL_START | room=%s | protocol=%s | session_id=%s | bridge=%s", + ctx.room.name, + protocol, + session_id, + bridge_identity or "-", + ) + log_flow_event( + call_logger, + "call_start", + room=ctx.room.name, + protocol=protocol, + session_id=session_id, + bridge=bridge_identity or "-", + ) + log_flow_event( + call_logger, + "agent_runtime_config", + min_interrupt_s=MIN_INTERRUPT_S, + discard_audio_if_uninterruptible=DISCARD_AUDIO_IF_UNINTERRUPTIBLE, + min_endpointing_delay=MIN_ENDPOINTING_DELAY, + max_endpointing_delay=MAX_ENDPOINTING_DELAY, + stt_min_audio_ms=os.getenv("STT_MIN_AUDIO_MS", "160"), + stt_min_dbfs=os.getenv("STT_MIN_DBFS", "-50.0"), + vad_config_source=vad_config_source, + vad_decision_logging=vad_logging_config["log_decisions"], + vad_activity_logging=vad_logging_config["log_activity"], + vad_activity_min_prob=vad_logging_config["activity_min_probability"], + vad_min_speech_duration=vad_config["min_speech_duration"], + vad_min_silence_duration=vad_config["min_silence_duration"], + vad_activation_threshold=vad_config["activation_threshold"], + vad_deactivation_threshold=vad_config["deactivation_threshold"], + vad_prefix_padding_duration=vad_config["prefix_padding_duration"], + agent_wait_timeout_retry_vad_threshold_enabled=( + agent_wait_timeout_retry_vad_threshold_enabled + ), + agent_wait_timeout_retry_vad_activation_threshold=( + AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD + ), + agent_wait_timeout_retry_vad_deactivation_threshold=( + agent_wait_timeout_retry_vad_deactivation_threshold + ), + inflight_wait_interval_s=REMOTE_AGENT_INFLIGHT_WAIT_INTERVAL_S, + inflight_wait_timeout_s=REMOTE_AGENT_INFLIGHT_WAIT_TIMEOUT_S, + inflight_wait_max_notices=REMOTE_AGENT_INFLIGHT_WAIT_MAX_NOTICES, + pre_backend_wait_notice_fast_on_vad_pause=( + agent_runtime_config["pre_backend_wait_notice_fast_on_vad_pause"] + ), + pre_backend_wait_notice_fast_on_vad_pause_source=( + pre_backend_wait_notice_fast_source + ), + inflight_wait_short_audio_dir=REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR, + inflight_wait_long_audio_dir=REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR, + inflight_wait_long_audio_path=REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_PATH, + ) + timeline.emit( + "call_start", + backend=resolve_agent_backend_name(call_config, os.getenv("AGENT_BACKEND", "remote_ws")), + bridge_identity=bridge_identity, + remote_agent=remote_agent_context.get("agent", ""), + timeline_file=str(timeline.path), + ) + bridge_gateway = BridgeGateway( + room=ctx.room, + bridge_identity=bridge_identity, + protocol=protocol, + timeline=timeline, + stress_test=bool(fake_agent_overrides.get("responses")), + ) + export_service = ExportService() + runtime_ref: dict[str, CallRuntime | None] = {"runtime": None} + + def _structured_interruption_flag() -> bool: + runtime = runtime_ref["runtime"] + return bool(runtime and runtime.structured_recebimento_interruption_flag()) + + stt_overrides = resolve_stt_overrides(call_config) + stt_provider = (stt_overrides.get("provider") or os.getenv("STT_PROVIDER", "internal_http")).strip().lower() + if stt_provider not in {"", "internal_http", "fake"}: + raise RuntimeError(f"Unsupported STT provider: {stt_provider}") + + if stt_provider == "fake": + stt = FakeSTT( + language=stt_overrides.get("language") or os.getenv("STT_LANG", "pt-BR"), + timeline=timeline, + structured_log_context=structured_context, + structured_logger=call_logger, + structured_interruption_flag=_structured_interruption_flag, + ) + else: + disable_vosk = str(stt_overrides.get("disable_vosk") or "").strip().lower() in { + "1", + "true", + "yes", + "on", + } + stt_cfg = InternalSTTConfig( + url=require_env("STT_URL", context="STT provider internal_http", timeline=timeline), + api_key=( + stt_overrides.get("api_key") + or stt_overrides.get("STT_KEY") + or os.getenv("STT_KEY", "") + or "unknow" + ), + language=stt_overrides.get("language", "portuguese") or os.getenv("STT_LANG", ""), + min_prob_single_word=float( + stt_overrides.get("min_prob_single_word") or os.getenv("STT_MIN_PROB_SINGLE_WORD", "0.10") + ), + initial_prompt=stt_overrides.get("initial_prompt") or os.getenv("STT_INITIAL_PROMPT", ""), + config_override=stt_overrides.get("config_override") or os.getenv("STT_CONFIG_OVERRIDE", ""), + ) + stt = InternalHTTPSTT( + stt_cfg, + client=http_client, + vosk=None if disable_vosk else vosk, + vosk_lock=None if disable_vosk else vosk_lock, + timeline=timeline, + structured_log_context=structured_context, + structured_logger=call_logger, + structured_interruption_flag=_structured_interruption_flag, + ) + + tts_overrides = resolve_tts_overrides(call_config) + tts_provider = (tts_overrides.get("provider") or os.getenv("TTS_PROVIDER", "elevenlabs")).strip().lower() + if tts_provider not in {"", "elevenlabs", "azure", "fake", "xai"}: + raise RuntimeError(f"Unsupported TTS provider: {tts_provider}") + + if tts_provider == "fake": + tts = livekit_tts.StreamAdapter(tts=FakeSessionTTS()) + elif tts_provider in {"", "elevenlabs"}: + tts = elevenlabs.TTS( + model=tts_overrides.get("model_id") or os.getenv("ELEVENLABS_MODEL_ID", ""), + voice_id=tts_overrides.get("voice_id") or os.getenv("ELEVENLABS_VOICE_ID", ""), + api_key=os.getenv("ELEVENLABS_API_KEY", ""), + voice_settings=elevenlabs.VoiceSettings( + stability=0.35, + speed=1.02, + similarity_boost=0.75, + use_speaker_boost=True, + style=0.7, + ), + language="pt", + ) + elif tts_provider == "xai": + xai_auth_method = (os.getenv("XAI_TTS_AUTH_METHOD", "API_KEY") or "API_KEY").strip().upper() + xai_auth_options: dict[str, str] = {"auth_method": xai_auth_method} + if xai_auth_method == "API_KEY": + xai_auth_options["api_key"] = require_env( + "XAI_API_KEY", context="TTS provider xai", timeline=timeline + ) + tts = OraclexAITTS( + voice=tts_overrides.get("voice_id") or os.getenv("XAI_TTS_VOICE", ""), + language=tts_overrides.get("language") or os.getenv("XAI_TTS_LANGUAGE", "pt-BR"), + initial_greeting_audio_cache=initial_greeting_audio_cache, + initial_greeting_agent=str(remote_agent_context.get("agent") or AGENT_NAME), + **xai_auth_options, + ) + else: + azure_tts_config, missing_keys = resolve_azure_speech_tts_config(tts_overrides) + for key in missing_keys: + if timeline is not None: + timeline.emit( + "config_missing", + key=key, + context="TTS provider azure", + ) + if missing_keys: + missing = ", ".join(missing_keys) + raise RuntimeError( + f"Missing required Azure Speech TTS config: {missing}. " + "Add the Azure Speech variables to .env.dev before starting the agent." + ) + + azure_impl = (os.getenv("AZURE_TTS_IMPLEMENTATION", "plugin") or "plugin").strip().lower() + if azure_impl not in {"plugin", "rest"}: + raise RuntimeError( + f"Unsupported AZURE_TTS_IMPLEMENTATION: {azure_impl}. " + "Use 'plugin' or 'rest'." + ) + + if azure_impl == "rest": + tts = AzureRESTTTS( + sample_rate=AZURE_SPEECH_TTS_SAMPLE_RATE, + **azure_tts_config, + ) + else: + try: + from livekit.plugins import azure as livekit_azure + except ModuleNotFoundError as exc: + raise RuntimeError( + "Azure Speech plugin is not installed. " + "Run make setup or reinstall requirements.txt before using TTS_PROVIDER=azure." + ) from exc + + tts = livekit_azure.TTS( + sample_rate=AZURE_SPEECH_TTS_SAMPLE_RATE, + **azure_tts_config, + ) + + if isinstance(tts, OraclexAITTS): + async def _prewarm_initial_tts_connection() -> None: + timeout_s = max(0.1, float(os.getenv("XAI_TTS_PREWARM_TIMEOUT_S", "3"))) + if timeline is not None: + timeline.emit("initial_tts_prewarm_started", timeout_ms=round(timeout_s * 1000)) + try: + await tts.prewarm_connection(timeout=timeout_s) + except Exception as exc: + logger.info("INITIAL_TTS_PREWARM_FAILED | error=%s", type(exc).__name__) + if timeline is not None: + timeline.emit("initial_tts_prewarm_failed", error=type(exc).__name__) + else: + if timeline is not None: + timeline.emit("initial_tts_prewarm_completed") + + create_task_logged(_prewarm_initial_tts_connection(), name="initial_tts_prewarm") + + vad_threshold_controller = ( + DynamicVADThresholdController( + vad, + config=DynamicVADThresholdConfig( + enabled=agent_wait_timeout_retry_vad_threshold_enabled, + baseline_activation_threshold=vad_config["activation_threshold"], + baseline_deactivation_threshold=vad_config["deactivation_threshold"], + retry_activation_threshold=AGENT_WAIT_TIMEOUT_RETRY_VAD_ACTIVATION_THRESHOLD, + retry_deactivation_threshold=agent_wait_timeout_retry_vad_deactivation_threshold, + ), + call_logger=call_logger, + timeline=timeline, + ) + if AGENT_WAIT_TIMEOUT_RETRY_VAD_THRESHOLD_ENABLED + else None + ) + + pre_backend_wait_notice_bridge = PreBackendWaitNoticeBridge() + # The deferred-interruption threshold is decided from END_OF_SPEECH events, + # so this wrapper is required even when verbose VAD logging is disabled. + wrap_vad = True + session_vad = ( + FlowLoggingVAD( + vad, + call_logger=call_logger, + min_interrupt_s=MIN_INTERRUPT_S, + vad_config=vad_config, + log_decisions=vad_logging_config["log_decisions"], + log_activity=vad_logging_config["log_activity"], + activity_min_prob=vad_logging_config["activity_min_probability"], + threshold_provider=( + vad_threshold_controller.current_thresholds + if vad_threshold_controller is not None + else None + ), + on_pre_backend_wait_pause=pre_backend_wait_notice_bridge.notify, + on_speech_end=pre_backend_wait_notice_bridge.notify_speech_end, + ) + if wrap_vad + else vad + ) + + session = AgentSession( + stt=stt, + tts=tts, + vad=session_vad, + discard_audio_if_uninterruptible=DISCARD_AUDIO_IF_UNINTERRUPTIBLE, + min_interruption_duration=MIN_INTERRUPT_S, + min_endpointing_delay=MIN_ENDPOINTING_DELAY, + max_endpointing_delay=MAX_ENDPOINTING_DELAY, + use_tts_aligned_transcript=True, + user_away_timeout=None, + ) + speech_service = SpeechService(session) + + agent = PipelineVoiceAgent( + elegibility=elegibility, + protocol=protocol, + intro=intro, + backend_name=resolve_agent_backend_name(call_config, os.getenv("AGENT_BACKEND", "remote_ws")), + remote_agent_context=remote_agent_context, + timeline=timeline, + ) + command_executor = RuntimeCommandExecutor( + agent=agent, + bridge_gateway=bridge_gateway, + export_service=export_service, + session=session, + speech_service=speech_service, + ) + interrupt_policy = InterruptPolicy() + idle_policy = IdlePolicy( + IdlePolicyConfig( + enabled=IDLE_NUDGE_ENABLED, + delay_s=IDLE_NUDGE_DELAY_S, + join_delay_s=IDLE_NUDGE_JOIN_DELAY_S, + close_delay_s=IDLE_NUDGE_CLOSE_DELAY_S, + max_tries=IDLE_NUDGE_MAX_TRIES, + end_reason=IDLE_NUDGE_END_REASON, + ) + ) + finalization_policy = FinalizationPolicy(call_end_grace_s=CALL_END_GRACE_S) + runtime = CallRuntime( + ctx=ctx, + session=session, + agent=agent, + command_executor=command_executor, + call_logger=call_logger, + protocol=protocol, + session_id=session_id, + bridge_identity=bridge_identity, + agent_starts_conversation=agent_starts_conversation, + nudge_text=nudge_text, + call_t0=call_t0, + extract_text_from_transcript=extract_text_from_transcript, + interrupt_policy=interrupt_policy, + idle_policy=idle_policy, + finalization_policy=finalization_policy, + create_task_logged=create_task_logged, + vad_threshold_controller=vad_threshold_controller, + initial_greeting_tts=tts, + config=RuntimeConfig( + call_end_grace_s=CALL_END_GRACE_S, + final_grace_s=FINAL_GRACE_S, + idle_nudge_delay_s=IDLE_NUDGE_DELAY_S, + idle_nudge_join_delay_s=IDLE_NUDGE_JOIN_DELAY_S, + idle_nudge_close_delay_s=IDLE_NUDGE_CLOSE_DELAY_S, + idle_nudge_max_tries=IDLE_NUDGE_MAX_TRIES, + idle_nudge_end_reason=IDLE_NUDGE_END_REASON, + inflight_backend_wait_interval_s=REMOTE_AGENT_INFLIGHT_WAIT_INTERVAL_S, + inflight_backend_wait_timeout_s=REMOTE_AGENT_INFLIGHT_WAIT_TIMEOUT_S, + inflight_backend_wait_max_notices=REMOTE_AGENT_INFLIGHT_WAIT_MAX_NOTICES, + inflight_backend_wait_text=REMOTE_AGENT_INFLIGHT_WAIT_TEXT, + inflight_backend_wait_short_audio_dir=REMOTE_AGENT_INFLIGHT_WAIT_SHORT_AUDIO_DIR, + inflight_backend_wait_long_audio_dir=REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_DIR, + deferred_interruption_min_audio_ms=agent_runtime_config[ + "deferred_interruption_min_audio_ms" + ], + deferred_interruption_enabled=agent_runtime_config[ + "deferred_interruption_enabled" + ], + deferred_interruption_stt_settle_timeout_s=DEFERRED_INTERRUPTION_STT_SETTLE_TIMEOUT_S, + deferred_interruption_user_turn_timeout_s=DEFERRED_INTERRUPTION_USER_TURN_TIMEOUT_S, + inflight_backend_wait_long_audio_path=REMOTE_AGENT_INFLIGHT_WAIT_LONG_AUDIO_PATH, + pre_backend_wait_notice_fast_on_vad_pause=( + agent_runtime_config["pre_backend_wait_notice_fast_on_vad_pause"] + ), + ), + logger=logger, + timeline=timeline, + structured_log_context=structured_context, + debug_event_publisher=( + bridge_gateway.publish_debug_event if debug_events_enabled else None + ), + ) + runtime_ref["runtime"] = runtime + empty_transcript_handler = getattr(stt, "set_empty_transcript_handler", None) + if callable(empty_transcript_handler): + empty_transcript_handler( + lambda: runtime.handle_empty_stt_final(source="stt_provider") + ) + stt_metrics_handler = getattr(stt, "set_metrics_handler", None) + if callable(stt_metrics_handler): + stt_metrics_handler(runtime.publish_stt_provider_metrics) + pre_backend_wait_notice_bridge.bind(runtime) + await runtime.run() + + +if __name__ == "__main__": + cli.run_app(server) diff --git a/src/app/livekit/policies/agent_finalization.py b/src/app/livekit/policies/agent_finalization.py new file mode 100644 index 0000000..ddb4abb --- /dev/null +++ b/src/app/livekit/policies/agent_finalization.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Mapping + + +AGENT_FINAL_RESULT_STOP_STATUS = { + "resolvido": "stop_resolvido_e_finalizado", + "nao_resolvido": "stop_nao_resolvido", + "resolvido_outros_assuntos": "stop_outro_assunto", + "outros_assuntos": "stop_outro_assunto", + "erro_falha_sistema": "stop_falha_sistema", + "erro_no_match": "stop_no_match", +} + + +@dataclass(frozen=True, slots=True) +class AgentFinalStop: + result_type: str + status: str + reason: str + phase: str = "in_session" + + +def normalize_agent_result_type(value: Any) -> str: + return str(value or "").strip().lower() + + +def stop_status_for_agent_result_type(result_type: Any) -> str: + return AGENT_FINAL_RESULT_STOP_STATUS.get(normalize_agent_result_type(result_type), "") + + +def final_stop_from_agent_result(result: Any) -> AgentFinalStop | None: + if not isinstance(result, Mapping): + return None + + result_type = normalize_agent_result_type(result.get("type")) + status = stop_status_for_agent_result_type(result_type) + if not status: + return None + + return AgentFinalStop( + result_type=result_type, + status=status, + reason=result_type, + ) diff --git a/src/app/livekit/policies/finalization_policy.py b/src/app/livekit/policies/finalization_policy.py new file mode 100644 index 0000000..8dd5c56 --- /dev/null +++ b/src/app/livekit/policies/finalization_policy.py @@ -0,0 +1,15 @@ +from __future__ import annotations + + +class FinalizationPolicy: + def __init__(self, *, call_end_grace_s: float) -> None: + self.call_end_grace_s = call_end_grace_s + + def should_skip(self, *, finalized: bool) -> bool: + return finalized + + def should_finalize_room_empty(self, *, remote_participants: int, finalized: bool) -> bool: + return remote_participants == 0 and not finalized + + def is_done_stage(self, stage: str) -> bool: + return (stage or "").upper() == "DONE" diff --git a/src/app/livekit/policies/idle_policy.py b/src/app/livekit/policies/idle_policy.py new file mode 100644 index 0000000..0a30476 --- /dev/null +++ b/src/app/livekit/policies/idle_policy.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional + + +@dataclass(frozen=True, slots=True) +class IdlePolicyConfig: + enabled: bool + delay_s: float + join_delay_s: float + close_delay_s: float + max_tries: int + end_reason: str + + +class IdlePolicy: + def __init__(self, config: IdlePolicyConfig) -> None: + self._config = config + + @property + def close_delay_s(self) -> float: + return self._config.close_delay_s + + @property + def max_tries(self) -> int: + return self._config.max_tries + + @property + def end_reason(self) -> str: + return self._config.end_reason or "no_user_response" + + @property + def enabled(self) -> bool: + return bool(self._config.enabled) + + @property + def join_delay_s(self) -> float: + return self._config.join_delay_s + + def resolve_delay(self, delay_s: Optional[float] = None) -> float: + return self._config.delay_s if delay_s is None else float(delay_s) + + def should_arm_nudge(self, *, delay_s: float, nudge_text: str) -> bool: + if not self.enabled: + return False + if delay_s <= 0: + return False + if not nudge_text: + return False + if self._config.max_tries <= 0: + return False + return True + + def should_fire_nudge( + self, + *, + token_is_current: bool, + seq_matches: bool, + speaking: bool, + gap_active: bool, + finalized: bool, + current_stage: str, + ) -> bool: + if not self.enabled: + return False + if not token_is_current: + return False + if not seq_matches: + return False + if speaking or gap_active: + return False + if finalized: + return False + if (current_stage or "").upper() == "DONE": + return False + return True + + def should_arm_close_before_fire(self, *, nudge_count: int) -> bool: + if not self.enabled: + return False + return nudge_count >= self._config.max_tries + + def should_fire_close( + self, + *, + token_is_current: bool, + seq_matches: bool, + finalized: bool, + current_stage: str, + ) -> bool: + if not self.enabled: + return False + if not token_is_current: + return False + if not seq_matches: + return False + if finalized: + return False + if (current_stage or "").upper() == "DONE": + return False + return True + + def should_arm_close_after_nudge(self, *, stage: str, nudge_count: int) -> bool: + if not self.enabled: + return False + return (stage or "").upper() == "IDLE_NUDGE" and nudge_count >= self._config.max_tries diff --git a/src/app/livekit/policies/interrupt_policy.py b/src/app/livekit/policies/interrupt_policy.py new file mode 100644 index 0000000..89bdd8b --- /dev/null +++ b/src/app/livekit/policies/interrupt_policy.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from typing import Iterable, Mapping + + +DEFAULT_STAGE_POLICY = { + "ARGUMENTATION": True, + "DATA_CONFIRMATION": False, + "FORMALIZATION": False, + "PRESENTATION": True, + "DONE": False, + "INTRO": True, + "IDLE_NUDGE": False, + "AGENT_WAIT_TIMEOUT_RETRY": True, +} + +DEFAULT_BACKCHANNEL_TOKENS = { + "uhum", + "aham", + "hm", + "hmm", + "ok", + "certo", + "tá", + "ta", + "sim", +} + + +class InterruptPolicy: + def __init__( + self, + *, + stage_policy: Mapping[str, bool] | None = None, + backchannel_tokens: Iterable[str] | None = None, + ) -> None: + self._stage_policy = { + str(key).strip().upper(): bool(value) + for key, value in (stage_policy or DEFAULT_STAGE_POLICY).items() + } + self._backchannel_tokens = { + str(token).strip().lower() + for token in (backchannel_tokens or DEFAULT_BACKCHANNEL_TOKENS) + if str(token).strip() + } + + def allow_stage(self, stage: str) -> bool: + key = (stage or "").strip().upper() + return bool(self._stage_policy.get(key, True)) + + def should_ignore_backchannel(self, *, speaking: bool, text: str) -> bool: + if not speaking: + return False + if len(text) <= 2: + return True + lowered = (text or "").strip().lower() + return lowered in self._backchannel_tokens diff --git a/src/app/livekit/runtime/call_runtime.py b/src/app/livekit/runtime/call_runtime.py new file mode 100644 index 0000000..fa72a76 --- /dev/null +++ b/src/app/livekit/runtime/call_runtime.py @@ -0,0 +1,6016 @@ +from __future__ import annotations + +import asyncio +import inspect +import json +import os +import random +import re +import time +import uuid +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable, Iterable, Mapping, Optional + +from livekit.agents import UserInputTranscribedEvent, room_io + +from app.livekit.adapters.agent_backend import BackendReply +from app.livekit.policies.agent_finalization import AgentFinalStop, final_stop_from_agent_result +from app.livekit.runtime.commands import ( + EndServiceOnce, + ExportSession, + ExtractSpokenText, + InterruptSpeech, + InjectIdleNudge, + NotifyBridgeDone, + NotifyBridgeStop, + RunPipelineInput, + SetPendingInterrupt, + SetPipelineInterruption, + SetPipelineProcessingInterruption, + StartSession, + StartSpeech, + WaitForSpeechPlayout, +) +from app.livekit.runtime.scheduler import TimerScheduler +from app.livekit.runtime.state import ( + CallState, + DeferredVADUtterance, + PendingUserTurn, + RuntimeConfig, +) +from app.livekit.runtime.wav_audio import wav_audio_frames, wav_duration_ms +from app.utils.logging import ( + EVENT_ENVIO_MSG, + EVENT_RECEBIMENTO_MSG, + error_message_from_resource, + event_latency_ms, + log_flow_event, + log_structured_event, +) +from app.utils.turn_ids import ( + clear_started_turn_message_id, + consume_transcribed_turn, + next_turn_message_id, + peek_started_turn_message_id, +) + +POST_USER_FINAL_SPEECH_GRACE_S = 0.35 +PRE_BACKEND_WAIT_NOTICE_GUARD_S = 0.15 +PROTECTED_SPEECH_POST_PLAYOUT_DROP_GRACE_S = 0.250 +USER_AUDIO_INPUT_SETUP_TIMEOUT_S = 15.0 +AGENT_WAIT_TIMEOUT_STOP_STATUS = "stop_silencio_longo" +TTS_FAILURE_DETAIL = "falha de reproducao do TTS" +IN_SESSION_STOP_STATUS_BY_RESOURCE = { + "agent_runtime": "stop_agent_runtime_unavailable", + "agent_backend": "stop_agent_backend_unavailable", + "stt": "stop_stt_unavailable", + "tts": "stop_tts_unavailable", + "bridge": "stop_bridge_failed", +} +INFLIGHT_BACKEND_WAIT_STAGE = "AGENT_BACKEND_WAIT" +AGENT_WAIT_TIMEOUT_RETRY_VAD_STAGES = {"AGENT_WAIT_TIMEOUT_RETRY"} +EMPTY_STT_FINAL_DEDUP_WINDOW_S = 1.0 +PENDING_VAD_STT_MAX_AGE_S = 15.0 +PENDING_VAD_STT_MAX_UTTERANCES = 32 +DEFERRED_INTERRUPTION_COMFORT_TEXT = "Um instante" +DEFERRED_INTERRUPTION_COMFORT_AUDIO_PATH = ( + Path(__file__).resolve().parents[1] + / "assets" + / "comfort" + / "interruption" + / "01.wav" +) + + +class InflightBackendWaitTimedOut(Exception): + pass + + +@dataclass(frozen=True, slots=True) +class InflightBackendWaitNoticeResult: + started: bool + interrupted_by_user: bool = False + + def __bool__(self) -> bool: + return self.started and not self.interrupted_by_user + + +@dataclass(frozen=True, slots=True) +class AgentWaitTimeoutConfig: + timeout_s: float + retry_messages: tuple[str, ...] + + +@dataclass(frozen=True, slots=True) +class _TTSPartialFailure: + detail: str + provider_synthesis_ms: int | None + pcm_duration_ms: int | None + + +def _http_status_details_from_exception(exc: Exception) -> tuple[int | str | None, str | None]: + response = getattr(exc, "response", None) + status = getattr(exc, "status_code", None) + if status is None: + status = getattr(exc, "status", None) + if status is None and response is not None: + status = getattr(response, "status_code", None) + if status in ("", "NA"): + status = None + + desc = None + if response is not None: + desc = str(getattr(response, "reason_phrase", "") or "").strip() or None + if desc is None: + desc = str(getattr(exc, "message", "") or "").strip() or None + if desc is None and status is not None: + desc = str(exc).strip() or type(exc).__name__ + return status, desc + + +class CallRuntime: + def __init__( + self, + *, + ctx: Any, + session: Any, + agent: Any, + command_executor: Any, + call_logger: Any, + protocol: str, + session_id: str, + bridge_identity: str, + agent_starts_conversation: bool, + nudge_text: str, + call_t0: float, + extract_text_from_transcript, + interrupt_policy: Any, + idle_policy: Any, + finalization_policy: Any, + create_task_logged: Callable[..., asyncio.Task], + config: RuntimeConfig, + logger: Any, + vad_threshold_controller: Any | None = None, + initial_greeting_tts: Any | None = None, + timeline: Any | None = None, + structured_log_context: Any | None = None, + debug_event_publisher: Callable[..., Any] | None = None, + ) -> None: + self._ctx = ctx + self._session = session + self._agent = agent + self._command_executor = command_executor + self._call_logger = call_logger + self._protocol = protocol + self._session_id = session_id + self._bridge_identity = bridge_identity + self._agent_starts_conversation = bool(agent_starts_conversation) + self._nudge_text = nudge_text + self._call_t0 = call_t0 + self._extract_text_from_transcript = extract_text_from_transcript + self._interrupt_policy = interrupt_policy + self._idle_policy = idle_policy + self._finalization_policy = finalization_policy + self._create_task_logged = create_task_logged + self._initial_greeting_tts = initial_greeting_tts + self._vad_threshold_controller = vad_threshold_controller + self._config = config + self._logger = logger + self._timeline = timeline + self._structured_log_context = structured_log_context + self._debug_event_publisher = debug_event_publisher + self._scheduler = TimerScheduler(create_task_logged=create_task_logged, logger=logger) + + self.state = CallState() + self._initial_agent_turn_requested = False + self._user_audio_input_enabled = True + self._user_audio_input_gate_released = False + self._say_stage_seq = 0 + self._last_vad_speech_end_at = 0.0 + self._last_vad_speech_duration_ms: int | None = None + + self.speaking = asyncio.Event() + self.say_lock = asyncio.Lock() + self.pending_lock = asyncio.Lock() + self.finalize_lock = asyncio.Lock() + self.finalized = asyncio.Event() + self._backend_push_task: asyncio.Task[Any] | None = None + self._terminal_stop_lock = asyncio.Lock() + self._terminal_stop_sent = False + self._last_interrupt_task: asyncio.Task[Any] | None = None + self._inflight_backend_activity_at = 0.0 + self._pre_backend_wait_notice_task: asyncio.Task[Any] | None = None + self._pre_backend_wait_notice_generation = 0 + self._pre_backend_wait_notice_reserved = False + self._agent_wait_timeout_config: AgentWaitTimeoutConfig | None = None + self._last_agent_wait_timeout_config: AgentWaitTimeoutConfig | None = None + self._last_agent_wait_timeout_user_seq = 0 + self._last_agent_wait_timeout_source = "" + self._last_agent_wait_timeout_message_id = "" + self._last_agent_message_id = "" + self._empty_stt_recovery_pending = False + self._agent_state_name = "" + self._agent_speaking_seq = 0 + self._agent_speaking_at = 0.0 + self._agent_speaking_event = asyncio.Event() + self._user_state_name = "listening" + self._user_not_speaking = asyncio.Event() + self._user_not_speaking.set() + self._tts_tffb_ms_by_speech_handle_id: dict[str, int] = {} + self._tts_duration_ms_by_speech_handle_id: dict[str, int] = {} + self._tts_audio_duration_ms_by_speech_handle_id: dict[str, int] = {} + self._tts_turn_health_by_segment_id: dict[str, tuple[int, int, int, int]] = {} + self._tts_turn_health_by_speech_handle_id: dict[str, tuple[int, int, int, int]] = {} + self._tts_metric_events_by_speech_handle_id: dict[str, asyncio.Event] = {} + self._retryable_tts_errors_by_speech_handle_id: dict[str, str] = {} + self._tts_partial_failures_by_speech_handle_id: dict[str, _TTSPartialFailure] = {} + self._structured_error_event_keys: set[tuple[str, str, str, str, str, str]] = set() + self._structured_interruption_pending = False + self._deferred_long_stt_settled = asyncio.Event() + self._deferred_long_stt_settled.set() + self._deferred_empty_final_latched_at = 0.0 + self._drop_next_user_final_context: dict[str, Any] | None = None + self._backend_processing_drop_context: dict[str, Any] | None = None + self._post_playout_drop_context: dict[str, Any] | None = None + self._post_playout_drop_expires_at = 0.0 + self._inflight_backend_wait_audio_candidates_cache = ( + self._inflight_backend_wait_audio_candidates_from_config(config) + ) + self._inflight_backend_wait_audio_duration_cache, self._inflight_backend_wait_audio_failures = ( + self._preload_wait_audio_paths( + self._unique_wait_audio_candidates( + self._inflight_backend_wait_audio_candidates_cache + ), + logger=logger, + ) + ) + self._inflight_backend_wait_audio_text_cache: dict[tuple[Path, str], str] = {} + self._preload_inflight_backend_wait_audio_texts() + subscribe_tts_event = getattr(self._initial_greeting_tts, "on", None) + if callable(subscribe_tts_event): + subscribe_tts_event("xai_tts_turn_timing", self._record_xai_tts_turn_timing) + subscribe_tts_event("xai_tts_turn_failed", self._record_xai_tts_turn_failed) + + async def execute(self, command: Any) -> Any: + return await self._command_executor.execute(command) + + def _emit_debug_event(self, event: str, **data: Any) -> None: + if self._debug_event_publisher is None: + return + try: + self._create_task_logged( + self._publish_debug_event(event, **data), + name=f"debug_event_{str(event).replace('.', '_')}", + ) + except Exception: + self._logger.debug("failed scheduling debug event %s", event, exc_info=True) + + async def _publish_debug_event(self, event: str, **data: Any) -> None: + if self._debug_event_publisher is None: + return + try: + result = self._debug_event_publisher(event, **data) + if inspect.isawaitable(result): + await result + except Exception: + self._logger.debug("failed publishing debug event %s", event, exc_info=True) + + def publish_stt_provider_metrics(self, fields: dict[str, Any]) -> None: + data = dict(fields) + outcome = str(data.pop("event", "completed") or "completed").strip().lower() + self._emit_debug_event(f"stt.sofya.{outcome}", **data) + + def execute_now(self, command: Any) -> Any: + return self._command_executor.execute_now(command) + + def _activate_agent_wait_timeout_retry_vad_threshold( + self, + *, + attempt: int, + reason: str = "agent_wait_timeout_retry", + ) -> None: + controller = self._vad_threshold_controller + activate = getattr(controller, "activate_agent_wait_timeout_retry", None) + if not callable(activate): + return + try: + activate(attempt=attempt, reason=reason) + except Exception: + self._logger.exception("[vad] failed activating agent wait timeout retry threshold") + + def _restore_agent_wait_timeout_retry_vad_threshold(self, *, reason: str) -> None: + controller = self._vad_threshold_controller + restore = getattr(controller, "restore", None) + if not callable(restore): + return + try: + restore(reason=reason) + except Exception: + self._logger.exception("[vad] failed restoring agent wait timeout retry threshold") + + @staticmethod + def _speech_handle_id(handle: Any) -> str: + return str(getattr(handle, "id", "") or getattr(handle, "_id", "") or "").strip() + + @staticmethod + def _seconds_to_ms(value: Any) -> int | None: + try: + seconds = float(value) + except (TypeError, ValueError): + return None + if seconds < 0: + return None + return max(0, round(seconds * 1000)) + + @classmethod + def _positive_seconds_to_ms(cls, value: Any) -> int | None: + ms = cls._seconds_to_ms(value) + if ms is None or ms <= 0: + return None + return ms + + @staticmethod + def _nonnegative_int(value: Any) -> int | None: + try: + parsed = int(value) + except (TypeError, ValueError): + return None + return max(0, parsed) + + def _record_xai_tts_turn_timing( + self, event: Any, *, publish_debug: bool = True + ) -> None: + """Keep provider-stream health data until its LiveKit metric arrives.""" + if not isinstance(event, dict): + return + segment_id = str(event.get("segment_id") or "").strip() + gap_ms = self._nonnegative_int(event.get("max_audio_delta_gap_ms")) + underrun_ms = self._nonnegative_int(event.get("max_playout_underrun_0ms")) + if not segment_id or gap_ms is None or underrun_ms is None: + return + underflow_count = self._nonnegative_int(event.get("xai_micro_underflows")) or 0 + avg_underflow_ms = self._nonnegative_int(event.get("xai_avg_underrun_ms")) or 0 + ( + previous_gap_ms, + previous_underrun_ms, + previous_underflow_count, + previous_underflow_total_ms, + ) = self._tts_turn_health_by_segment_id.get(segment_id, (0, 0, 0, 0)) + self._tts_turn_health_by_segment_id[segment_id] = ( + max(previous_gap_ms, gap_ms), + max(previous_underrun_ms, underrun_ms), + previous_underflow_count + underflow_count, + previous_underflow_total_ms + underflow_count * avg_underflow_ms, + ) + if publish_debug: + self._emit_debug_event("tts.xai.completed", **event) + + def _record_xai_tts_turn_failed(self, event: Any) -> None: + """Bind a provider partial failure to the active speech immediately. + + LiveKit does not emit ``metrics_collected`` when a streamed response + fails after PCM was released. This event is the authoritative signal + for the structured error path in that case. + """ + if not isinstance(event, dict): + return + segment_id = str(event.get("segment_id") or "").strip() + reason = str(event.get("reason") or "provider_partial_failure").strip() + gap_ms = self._nonnegative_int(event.get("max_audio_delta_gap_ms")) + underrun_ms = self._nonnegative_int(event.get("max_playout_underrun_0ms")) + if not segment_id or gap_ms is None or underrun_ms is None: + return + + self._record_xai_tts_turn_timing(event, publish_debug=False) + self._emit_debug_event("tts.xai.failed", **event) + handle = self.state.current_speech.handle + speech_handle_id = self._speech_handle_id(handle) + if not speech_handle_id: + log_flow_event( + self._call_logger, + "tts_partial_failure_without_active_handle", + segment_id=segment_id, + reason=reason, + ) + return + + health = self._tts_turn_health_by_segment_id.get(segment_id) + if health is not None: + self._tts_turn_health_by_speech_handle_id[speech_handle_id] = health + underflow_count = self._nonnegative_int(event.get("xai_micro_underflows")) or 0 + avg_underflow_ms = self._nonnegative_int(event.get("xai_avg_underrun_ms")) or 0 + provider_synthesis_ms = self._nonnegative_int(event.get("provider_synthesis_ms")) + pcm_duration_ms = self._nonnegative_int(event.get("pcm_duration_ms")) + detail = ( + f"xAI TTS partial audio failure: {reason}; " + f"xai_underrun_estimado_ms={underrun_ms}; " + f"xai_micro_underflows={underflow_count}; " + f"xai_avg_underrun_ms={avg_underflow_ms}; " + f"max_audio_delta_gap_ms={gap_ms}" + ) + self._tts_partial_failures_by_speech_handle_id[speech_handle_id] = ( + _TTSPartialFailure( + detail=detail, + provider_synthesis_ms=provider_synthesis_ms, + pcm_duration_ms=pcm_duration_ms, + ) + ) + log_flow_event( + self._call_logger, + "tts_partial_failure_observed", + stage=self.state.current_speech.stage or self.state.current_stage, + speech_id=self.state.current_speech.speech_id, + speech_handle_id=speech_handle_id, + segment_id=segment_id, + reason=reason, + max_gap_ms=gap_ms, + max_underrun_0ms=underrun_ms, + underflow_count=underflow_count, + avg_underflow_ms=avg_underflow_ms, + ) + + def _pop_tts_partial_failure(self, handle: Any) -> _TTSPartialFailure | None: + speech_handle_id = self._speech_handle_id(handle) + if not speech_handle_id: + return None + return self._tts_partial_failures_by_speech_handle_id.pop( + speech_handle_id, None + ) + + def _resolve_tts_turn_health_metrics( + self, handle: Any + ) -> tuple[int | None, int | None, int | None, int | None]: + speech_handle_id = self._speech_handle_id(handle) + if not speech_handle_id: + return None, None, None, None + health = self._tts_turn_health_by_speech_handle_id.get(speech_handle_id) + if health is None: + return None, None, None, None + gap_ms, max_underrun_ms, underflow_count, underflow_total_ms = health + return ( + gap_ms, + max_underrun_ms, + underflow_count, + round(underflow_total_ms / underflow_count) if underflow_count else 0, + ) + + def _record_tts_metric(self, ev: Any) -> None: + metrics = getattr(ev, "metrics", ev) + if str(getattr(metrics, "type", "") or "") != "tts_metrics": + return + + speech_handle_id = str(getattr(metrics, "speech_id", "") or "").strip() + if not speech_handle_id: + speech_handle_id = self._speech_handle_id(self.state.current_speech.handle) + if not speech_handle_id: + return + + segment_id = str(getattr(metrics, "segment_id", "") or "").strip() + segment_health = self._tts_turn_health_by_segment_id.get(segment_id) + if segment_health is not None: + self._tts_turn_health_by_speech_handle_id.setdefault( + speech_handle_id, segment_health + ) + duration_ms = self._seconds_to_ms(getattr(metrics, "duration", None)) + if duration_ms is not None and duration_ms > 0: + self._tts_duration_ms_by_speech_handle_id.setdefault( + speech_handle_id, + duration_ms, + ) + + audio_ms = self._positive_seconds_to_ms(getattr(metrics, "audio_duration", None)) + if audio_ms is not None: + self._tts_audio_duration_ms_by_speech_handle_id.setdefault( + speech_handle_id, + audio_ms, + ) + + ttfb_ms = self._positive_seconds_to_ms(getattr(metrics, "ttfb", None)) + if ttfb_ms is None: + metric_event = self._tts_metric_events_by_speech_handle_id.get(speech_handle_id) + if metric_event is not None: + metric_event.set() + log_flow_event( + self._call_logger, + "tts_metric_invalid", + speech_handle_id=speech_handle_id, + tffb_seconds=getattr(metrics, "ttfb", None), + duration_ms=duration_ms, + audio_ms=audio_ms, + provider=str(getattr(metrics, "label", "") or ""), + ) + return + + self._tts_tffb_ms_by_speech_handle_id.setdefault(speech_handle_id, ttfb_ms) + metric_event = self._tts_metric_events_by_speech_handle_id.get(speech_handle_id) + if metric_event is not None: + metric_event.set() + log_flow_event( + self._call_logger, + "tts_metric", + speech_handle_id=speech_handle_id, + tffb_ms=ttfb_ms, + duration_ms=duration_ms, + audio_ms=audio_ms, + provider=str(getattr(metrics, "label", "") or ""), + ) + self._emit_debug_event( + "tts.completed", + stage=self.state.current_speech.stage or self.state.current_stage, + speech_id=speech_handle_id, + duration_ms=duration_ms, + ttfb_ms=ttfb_ms, + audio_duration_ms=audio_ms, + provider=str(getattr(metrics, "label", "") or ""), + max_audio_delta_gap_ms=( + segment_health[0] if segment_health is not None else None + ), + max_playout_underrun_ms=( + segment_health[1] if segment_health is not None else None + ), + underflow_count=( + segment_health[2] if segment_health is not None else None + ), + avg_underflow_ms=( + round(segment_health[3] / segment_health[2]) + if segment_health is not None and segment_health[2] + else (0 if segment_health is not None else None) + ), + ) + + async def _wait_for_tts_metric_event(self, speech_handle_id: str) -> None: + wait_s = self._tts_metric_wait_s() + if wait_s <= 0: + return + + metric_event = self._tts_metric_events_by_speech_handle_id.setdefault( + speech_handle_id, + asyncio.Event(), + ) + try: + await asyncio.wait_for(metric_event.wait(), timeout=wait_s) + except asyncio.TimeoutError: + pass + finally: + if self._tts_metric_events_by_speech_handle_id.get(speech_handle_id) is metric_event: + self._tts_metric_events_by_speech_handle_id.pop(speech_handle_id, None) + + def _has_tts_metric_for_speech_handle_id(self, speech_handle_id: str) -> bool: + return any( + metrics.get(speech_handle_id) is not None + for metrics in ( + self._tts_tffb_ms_by_speech_handle_id, + self._tts_duration_ms_by_speech_handle_id, + self._tts_audio_duration_ms_by_speech_handle_id, + ) + ) + + @staticmethod + def _tts_metric_wait_s() -> float: + try: + wait_ms = int(str(os.getenv("TTS_TTFB_METRIC_WAIT_MS", "250") or "250")) + except ValueError: + wait_ms = 250 + return max(0.0, wait_ms / 1000.0) + + @staticmethod + def _float_env_s(env_name: str, default: float) -> float: + try: + timeout_s = float(str(os.getenv(env_name, str(default)) or str(default))) + except ValueError: + timeout_s = default + return max(0.0, timeout_s) + + @classmethod + def _tts_first_frame_timeout_s(cls) -> float: + if os.getenv("TTS_FIRST_FRAME_TIMEOUT_S") is not None: + return cls._float_env_s("TTS_FIRST_FRAME_TIMEOUT_S", 3.0) + return cls._float_env_s("TTS_EMPTY_FRAME_RETRY_TIMEOUT_S", 3.0) + + @classmethod + def _tts_connection_start_grace_s(cls) -> float: + return cls._float_env_s("TTS_CONNECTION_START_GRACE_TIMEOUT_S", 3.0) + + @classmethod + def _tts_playout_start_timeout_s(cls) -> float: + if os.getenv("TTS_PLAYOUT_START_TIMEOUT_S") is not None: + return cls._float_env_s("TTS_PLAYOUT_START_TIMEOUT_S", 6.0) + return cls._tts_first_frame_timeout_s() + cls._tts_connection_start_grace_s() + + @classmethod + def _tts_empty_frame_retry_timeout_s(cls) -> float: + return cls._tts_playout_start_timeout_s() + + @staticmethod + def _exception_chain_text(exc: Exception) -> str: + parts: list[str] = [] + current: BaseException | None = exc + seen: set[int] = set() + while current is not None and id(current) not in seen: + seen.add(id(current)) + parts.append(type(current).__name__) + parts.append(str(current)) + current = current.__cause__ or current.__context__ + return " ".join(parts).lower() + + @staticmethod + def _tts_failure_metric(detail: str | None, key: str) -> int | None: + match = re.search(rf"(?:^|;\s*){re.escape(key)}=(\d+)(?:;|$)", str(detail or "")) + return int(match.group(1)) if match else None + + @classmethod + def _is_retryable_tts_playout_error(cls, exc: Exception) -> bool: + return cls._is_retryable_tts_error_text(cls._exception_chain_text(exc)) + + @staticmethod + def _retryable_tts_error_markers() -> tuple[str, ...]: + return ( + "audio frame gap", + "no audio frames", + "empty frame", + "apitimeouterror", + "request timed out", + "xai tts partial audio failure", + "xai tts total timeout", + "xai tts failed before audio", + "failed to connect to xai tts", + ) + + @classmethod + def _is_retryable_tts_error_text(cls, details: str) -> bool: + details = str(details or "").lower() + return any(marker in details for marker in cls._retryable_tts_error_markers()) + + @classmethod + def _retryable_tts_error_detail(cls, error: Any) -> str | None: + underlying = getattr(error, "error", error) + if isinstance(underlying, BaseException): + details = cls._exception_chain_text(underlying) + else: + details = f"{type(underlying).__name__} {underlying}".lower() + if not cls._is_retryable_tts_error_text(details): + return None + return details[:500] + + def _record_retryable_tts_session_error(self, error: Any) -> None: + error_type = str(getattr(error, "type", "") or "").strip().lower() + if error_type != "tts_error": + return + + detail = self._retryable_tts_error_detail(error) + if detail is None: + return + + speech_handle_id = self._speech_handle_id(self.state.current_speech.handle) + if not speech_handle_id: + return + + self._retryable_tts_errors_by_speech_handle_id[speech_handle_id] = detail + log_flow_event( + self._call_logger, + "tts_retryable_error_observed", + stage=self.state.current_speech.stage or self.state.current_stage, + speech_id=self.state.current_speech.speech_id, + speech_handle_id=speech_handle_id, + recoverable=bool(getattr(error, "recoverable", False)), + detail=detail, + ) + + def _pop_retryable_tts_error_detail(self, handle: Any) -> str | None: + speech_handle_id = self._speech_handle_id(handle) + if not speech_handle_id: + return None + return self._retryable_tts_errors_by_speech_handle_id.pop(speech_handle_id, None) + + def _agent_speaking_seen_after( + self, + *, + started_monotonic: float, + seq_at_start: int, + ) -> bool: + session_state = str(getattr(self._session, "agent_state", "") or "").lower() + session_state_name = session_state.rsplit(".", 1)[-1] + if session_state_name == "speaking" and self._agent_speaking_at >= started_monotonic: + return True + return ( + self._agent_speaking_seq > seq_at_start + and self._agent_speaking_at >= started_monotonic + ) + + async def _wait_for_agent_speaking_after( + self, + *, + started_monotonic: float, + seq_at_start: int, + timeout_s: float, + ) -> bool: + deadline = time.monotonic() + timeout_s + while True: + if self._agent_speaking_seen_after( + started_monotonic=started_monotonic, + seq_at_start=seq_at_start, + ): + return True + + remaining_s = deadline - time.monotonic() + if remaining_s <= 0: + return False + + self._agent_speaking_event.clear() + try: + await asyncio.wait_for( + self._agent_speaking_event.wait(), + timeout=min(remaining_s, 0.05), + ) + except asyncio.TimeoutError: + pass + + async def _wait_for_tts_playout_or_empty_frame_retry( + self, + handle: Any, + *, + started_monotonic: float, + agent_speaking_seq_at_start: int, + timeout_s: float, + stage: str, + message_id: str, + ) -> bool: + if timeout_s <= 0: + await self.execute(WaitForSpeechPlayout(handle)) + return False + + playout_task = asyncio.create_task(self.execute(WaitForSpeechPlayout(handle))) + speaking_task = asyncio.create_task( + self._wait_for_agent_speaking_after( + started_monotonic=started_monotonic, + seq_at_start=agent_speaking_seq_at_start, + timeout_s=timeout_s, + ) + ) + try: + done, _ = await asyncio.wait( + {playout_task, speaking_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + if speaking_task in done: + first_audio_seen = await speaking_task + if first_audio_seen: + await playout_task + return False + + if playout_task.done() and playout_task.exception() is not None: + await playout_task + + elif playout_task in done: + await playout_task + return False + + speech_handle_id = self._speech_handle_id(handle) + log_flow_event( + self._call_logger, + "tts_empty_frame_retry", + stage=stage, + message_id=message_id, + speech_handle_id=speech_handle_id, + timeout_ms=round(timeout_s * 1000), + timeout_kind="playout_start", + first_frame_timeout_ms=round(self._tts_first_frame_timeout_s() * 1000), + connection_grace_ms=round(self._tts_connection_start_grace_s() * 1000), + action="resend_text", + ) + try: + await self.execute(InterruptSpeech(handle, force=True)) + except Exception as exc: + log_flow_event( + self._call_logger, + "tts_empty_frame_retry_interrupt_error", + stage=stage, + message_id=message_id, + speech_handle_id=speech_handle_id, + error=type(exc).__name__, + ) + + playout_task.cancel() + await asyncio.gather(playout_task, return_exceptions=True) + return True + finally: + if not speaking_task.done(): + speaking_task.cancel() + await asyncio.gather(speaking_task, return_exceptions=True) + + async def _resolve_tts_tffb_ms( + self, + handle: Any, + ) -> int | None: + speech_handle_id = self._speech_handle_id(handle) + if not speech_handle_id: + return None + + ttfb_ms = self._tts_tffb_ms_by_speech_handle_id.get(speech_handle_id) + if ttfb_ms is not None and ttfb_ms > 0: + return ttfb_ms + + if self._has_tts_metric_for_speech_handle_id(speech_handle_id): + return None + + await self._wait_for_tts_metric_event(speech_handle_id) + ttfb_ms = self._tts_tffb_ms_by_speech_handle_id.get(speech_handle_id) + if ttfb_ms is not None and ttfb_ms > 0: + return ttfb_ms + + return None + + def _resolve_tts_provider_duration_ms( + self, + handle: Any, + ) -> int | None: + speech_handle_id = self._speech_handle_id(handle) + if not speech_handle_id: + return None + + duration_ms = self._tts_duration_ms_by_speech_handle_id.get(speech_handle_id) + if duration_ms is not None and duration_ms > 0: + return duration_ms + + return None + + def _resolve_tts_structured_metrics( + self, + handle: Any, + ) -> tuple[int | None, int | None, int | None, int | None, str]: + speech_handle_id = self._speech_handle_id(handle) + if not speech_handle_id: + return None, None, None, None, "livekit_tts_metrics_missing" + + ttfb_ms = self._tts_tffb_ms_by_speech_handle_id.get(speech_handle_id) + metric_duration_ms = self._tts_duration_ms_by_speech_handle_id.get(speech_handle_id) + audio_duration_ms = self._tts_audio_duration_ms_by_speech_handle_id.get(speech_handle_id) + has_livekit_metric = any( + value is not None and value > 0 + for value in (ttfb_ms, metric_duration_ms, audio_duration_ms) + ) + metrics_source = ( + "livekit_tts_metrics" if has_livekit_metric else "livekit_tts_metrics_missing" + ) + total_ms = metric_duration_ms + return total_ms, ttfb_ms, audio_duration_ms, metric_duration_ms, metrics_source + + def _resolve_tts_audio_duration_ms( + self, + handle: Any, + ) -> int | None: + speech_handle_id = self._speech_handle_id(handle) + if speech_handle_id: + audio_ms = self._tts_audio_duration_ms_by_speech_handle_id.get(speech_handle_id) + if audio_ms is not None and audio_ms > 0: + return audio_ms + + return None + + def _mark_structured_interruption(self, *, reason: str = "") -> None: + self._structured_interruption_pending = True + log_flow_event( + self._call_logger, + "structured_interruption_marked", + reason=reason, + stage=self.state.current_speech.stage or self.state.current_stage, + ) + + def _clear_structured_interruption(self) -> None: + self._structured_interruption_pending = False + + def structured_recebimento_interruption_flag(self) -> bool: + return bool( + self._structured_interruption_pending + or (self.speaking.is_set() and self._current_speech_allows_interruption()) + ) + + def reset_gap_state(self) -> None: + self.state.gap.active = False + self.state.gap.guard_seq = 0 + self.state.gap.stage = "" + self.state.gap.text = "" + self.state.gap.speech_id = "" + self.state.gap.cancelled = False + + @staticmethod + def _metadata_bool(metadata: Any, key: str) -> Optional[bool]: + if not isinstance(metadata, Mapping) or key not in metadata: + return None + + value = metadata.get(key) + if isinstance(value, bool): + return value + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"1", "true", "yes", "on"}: + return True + if normalized in {"0", "false", "no", "off"}: + return False + return None + + @staticmethod + def _metadata_str(metadata: Any, key: str) -> str: + if not isinstance(metadata, Mapping): + return "" + return str(metadata.get(key) or "").strip() + + @classmethod + def _reply_expects_user_response(cls, reply: BackendReply) -> Optional[bool]: + return cls._metadata_bool(reply.metadata, "expects_user_response") + + @classmethod + def _reply_drops_user_input_while_speaking(cls, reply: BackendReply) -> bool: + return bool(cls._metadata_bool(reply.metadata, "drop_user_input_while_speaking")) + + @classmethod + def _reply_agent_message_type(cls, reply: BackendReply) -> str: + return cls._metadata_str(reply.metadata, "agent_message_type") + + @classmethod + def _reply_agent_result_type(cls, reply: BackendReply) -> str: + return cls._metadata_str(reply.metadata, "agent_result_type") + + def _protected_speech_drop_context_from_values( + self, + *, + reason: str, + stage: str, + allow_interruptions: bool, + drop_user_input_while_speaking: bool, + agent_message_type: str, + agent_result_type: str, + ) -> dict[str, Any] | None: + stage = stage or self.state.current_stage + protected_by_metadata = bool(drop_user_input_while_speaking) + protected_presentation = ( + str(stage or "").strip().upper() == "PRESENTATION" + and not allow_interruptions + ) + if not protected_by_metadata and not protected_presentation: + return None + + agent_message_type = str(agent_message_type or "").strip() + if not agent_message_type and protected_presentation: + agent_message_type = "ready" if self.state.user_final_seq == 0 else "presentation" + + return { + "reason": reason, + "stage": stage, + "agent_message_type": agent_message_type, + "agent_result_type": str(agent_result_type or "").strip(), + } + + def _current_speech_drop_context(self, *, reason: str) -> dict[str, Any] | None: + current = self.state.current_speech + if not self.speaking.is_set(): + return None + return self._protected_speech_drop_context_from_values( + reason=reason, + stage=current.stage or self.state.current_stage, + allow_interruptions=current.allow_interruptions, + drop_user_input_while_speaking=current.drop_user_input_while_speaking, + agent_message_type=current.agent_message_type, + agent_result_type=current.agent_result_type, + ) + + def _mark_drop_next_user_final(self, context: Mapping[str, Any]) -> bool: + self._drop_next_user_final_context = dict(context) + self._call_logger.info( + "USER_INPUT_DROP_ARMED | reason=%s | stage=%s | agent_message_type=%s | agent_result_type=%s", + context.get("reason") or "-", + context.get("stage") or "-", + context.get("agent_message_type") or "-", + context.get("agent_result_type") or "-", + ) + log_flow_event( + self._call_logger, + "user_input_drop_armed", + reason=context.get("reason") or "", + stage=context.get("stage") or "", + agent_message_type=context.get("agent_message_type") or "", + agent_result_type=context.get("agent_result_type") or "", + ) + if self._timeline is not None: + self._timeline.emit("user_input_drop_armed", **context) + return True + + def _mark_drop_next_user_final_from_current_speech(self, *, reason: str) -> bool: + context = self._current_speech_drop_context(reason=reason) + if context is None: + return False + return self._mark_drop_next_user_final(context) + + def _active_post_playout_drop_context(self, *, reason: str) -> dict[str, Any] | None: + context = self._post_playout_drop_context + if context is None: + return None + if time.monotonic() >= self._post_playout_drop_expires_at: + self._clear_post_playout_drop_window(reason="expired") + return None + active_context = dict(context) + active_context["reason"] = reason + return active_context + + def _mark_drop_next_user_final_from_post_playout_window(self, *, reason: str) -> bool: + context = self._active_post_playout_drop_context(reason=reason) + if context is None: + return False + return self._mark_drop_next_user_final(context) + + def _arm_post_playout_drop_window(self, context: Mapping[str, Any]) -> None: + grace_s = PROTECTED_SPEECH_POST_PLAYOUT_DROP_GRACE_S + if grace_s <= 0: + return + self._post_playout_drop_context = dict(context) + self._post_playout_drop_expires_at = time.monotonic() + grace_s + self._call_logger.info( + "USER_INPUT_DROP_GRACE_ARMED | reason=%s | stage=%s | agent_message_type=%s | agent_result_type=%s | grace_ms=%s", + context.get("reason") or "-", + context.get("stage") or "-", + context.get("agent_message_type") or "-", + context.get("agent_result_type") or "-", + round(grace_s * 1000), + ) + log_flow_event( + self._call_logger, + "user_input_drop_grace_armed", + reason=context.get("reason") or "", + stage=context.get("stage") or "", + agent_message_type=context.get("agent_message_type") or "", + agent_result_type=context.get("agent_result_type") or "", + grace_ms=round(grace_s * 1000), + ) + if self._timeline is not None: + self._timeline.emit( + "user_input_drop_grace_armed", + **dict(context), + grace_ms=round(grace_s * 1000), + ) + + def _clear_post_playout_drop_window(self, *, reason: str = "clear") -> None: + if self._post_playout_drop_context is None: + return + self._post_playout_drop_context = None + self._post_playout_drop_expires_at = 0.0 + log_flow_event( + self._call_logger, + "user_input_drop_grace_cleared", + reason=reason, + ) + + def _clear_pending_user_input_drop(self) -> None: + self._drop_next_user_final_context = None + + def _backend_processing_drop_context_from_reply( + self, + reply: BackendReply, + *, + reason: str, + ) -> dict[str, Any]: + return { + "reason": reason, + "stage": reply.stage or self.state.current_stage, + "agent_message_type": self._reply_agent_message_type(reply), + "agent_result_type": self._reply_agent_result_type(reply), + } + + def _update_backend_processing_drop_from_reply( + self, + reply: BackendReply, + *, + source: str, + ) -> None: + if ( + self._reply_agent_message_type(reply) == "feedback" + or self._reply_agent_result_type(reply) == "feedback" + ): + self._backend_processing_drop_context = self._backend_processing_drop_context_from_reply( + reply, + reason=f"backend_processing_feedback:{source}", + ) + return + + self._backend_processing_drop_context = None + + def _drop_context_for_user_final(self) -> dict[str, Any] | None: + if self.finalized.is_set(): + return { + "reason": "finalized", + "stage": self.state.current_stage, + "agent_message_type": "", + "agent_result_type": "", + } + + context = self._current_speech_drop_context(reason="protected_speech_active") + if context is not None: + self._clear_pending_user_input_drop() + return context + + if self._backend_processing_drop_context is not None: + return dict(self._backend_processing_drop_context) + + if self._drop_next_user_final_context is not None: + context = dict(self._drop_next_user_final_context) + self._clear_pending_user_input_drop() + return context + + context = self._active_post_playout_drop_context( + reason="protected_speech_post_playout_grace" + ) + if context is not None: + self._clear_post_playout_drop_window(reason="user_final_dropped") + return context + + return None + + def _log_user_input_dropped( + self, + *, + context: Mapping[str, Any], + transcription: str, + text: str, + ) -> None: + reason = str(context.get("reason") or "protected_speech").strip() + stage = str(context.get("stage") or self.state.current_stage or "").strip() + agent_message_type = str(context.get("agent_message_type") or "").strip() + agent_result_type = str(context.get("agent_result_type") or "").strip() + self._call_logger.info( + "USER_INPUT_DROPPED | reason=%s | stage=%s | agent_message_type=%s | agent_result_type=%s | text=%r", + reason or "-", + stage or "-", + agent_message_type or "-", + agent_result_type or "-", + text, + ) + log_flow_event( + self._call_logger, + "user_input_dropped", + reason=reason, + stage=stage, + agent_message_type=agent_message_type, + agent_result_type=agent_result_type, + text=text, + transcript_len=len(transcription or ""), + ) + if self._timeline is not None: + self._timeline.emit( + "user_input_dropped", + reason=reason, + stage=stage, + agent_message_type=agent_message_type, + agent_result_type=agent_result_type, + text=text, + transcript_len=len(transcription or ""), + ) + + @staticmethod + def _speech_id_from_reply(reply: BackendReply) -> str: + metadata = reply.metadata + if not isinstance(metadata, Mapping): + return "" + return str(metadata.get("speech_id") or "").strip() + + @staticmethod + def _message_id_from_reply(reply: BackendReply) -> str: + metadata = reply.metadata + if not isinstance(metadata, Mapping): + return "" + + for key in ("message_id", "messageId"): + value = metadata.get(key) + if value not in (None, ""): + return str(value).strip() + + payload = metadata.get("payload") + if isinstance(payload, Mapping): + for key in ("message_id", "messageId"): + value = payload.get(key) + if value not in (None, ""): + return str(value).strip() + + return "" + + @staticmethod + def _backend_reply_with_message_id(reply: BackendReply, message_id: str) -> BackendReply: + message_id = str(message_id or "").strip() + if not message_id: + return reply + + metadata = dict(reply.metadata) if isinstance(reply.metadata, Mapping) else {} + metadata["message_id"] = message_id + return BackendReply( + stage=reply.stage, + text=reply.text, + done=reply.done, + export_payload=reply.export_payload, + metadata=metadata, + ) + + @staticmethod + def _payload_with_message_id(payload: Any, message_id: str, *, text: str = "") -> Any: + message_id = str(message_id or "").strip() + if not message_id: + return payload + + if isinstance(payload, Mapping): + enriched = dict(payload) + enriched["message_id"] = message_id + if text and not any( + key in enriched + for key in ("text", "transcript", "utterance", "message", "content") + ): + enriched["text"] = text + return enriched + + return { + "text": str(text or payload or "").strip(), + "message_id": message_id, + } + + def _consume_turn_message_id(self, transcription: str, text: str) -> str: + if not text: + return "" + + turn = consume_transcribed_turn( + self._structured_log_context, + transcription=transcription, + text=text, + ) + if turn is None: + return "" + clear_started_turn_message_id( + self._structured_log_context, + message_id=turn.message_id, + ) + return turn.message_id + + @staticmethod + def _combined_interruption_transcription(turns: list[PendingUserTurn]) -> tuple[str, str]: + combined_text = ". ".join( + turn.text.strip().rstrip(".") for turn in turns if turn.text.strip() + ).strip() + if not turns: + return "", combined_text + try: + payload = json.loads(turns[0].transcription) + except Exception: + return combined_text, combined_text + if not isinstance(payload, dict): + return combined_text, combined_text + data = payload.get("data") + if isinstance(data, dict): + payload = dict(payload) + payload["data"] = dict(data) + payload["data"]["text"] = combined_text + payload["data"].pop("words", None) + elif "text" in payload: + payload = dict(payload) + payload["text"] = combined_text + else: + payload = dict(payload) + payload["text"] = combined_text + return json.dumps(payload, ensure_ascii=False), combined_text + + def _prune_pending_vad_utterances(self) -> None: + deferred = self.state.deferred_interruption + now = time.monotonic() + queues = ( + ("pre_backend", deferred.pre_backend_vad_utterances), + ("backend", deferred.pending_vad_utterances), + ) + for queue_name, queue in queues: + expired = 0 + capacity_dropped = 0 + while ( + queue + and queue[0].ended_at + and now - queue[0].ended_at > PENDING_VAD_STT_MAX_AGE_S + ): + utterance = queue.popleft() + expired += 1 + if queue_name == "backend" and utterance.is_long: + deferred.pending_long_stt_finals = max( + 0, deferred.pending_long_stt_finals - 1 + ) + while len(queue) > PENDING_VAD_STT_MAX_UTTERANCES: + utterance = queue.popleft() + capacity_dropped += 1 + if queue_name == "backend" and utterance.is_long: + deferred.pending_long_stt_finals = max( + 0, deferred.pending_long_stt_finals - 1 + ) + if expired: + log_flow_event( + self._call_logger, + "deferred_interruption_vad_expired", + queue=queue_name, + discarded_utterances=expired, + max_age_s=PENDING_VAD_STT_MAX_AGE_S, + ) + if capacity_dropped: + log_flow_event( + self._call_logger, + "deferred_interruption_vad_capacity_shed", + queue=queue_name, + discarded_utterances=capacity_dropped, + max_utterances=PENDING_VAD_STT_MAX_UTTERANCES, + ) + if not deferred.pending_long_stt_finals: + self._deferred_long_stt_settled.set() + + def _promote_pre_backend_vad_utterances(self) -> None: + deferred = self.state.deferred_interruption + self._prune_pending_vad_utterances() + if not deferred.pre_backend_vad_utterances: + return + promoted = list(deferred.pre_backend_vad_utterances) + deferred.pre_backend_vad_utterances.clear() + deferred.pending_vad_utterances.extend(promoted) + long_count = sum(1 for utterance in promoted if utterance.is_long) + if not long_count: + return + if long_count: + deferred.pending_long_stt_finals += long_count + self._deferred_long_stt_settled.clear() + if not deferred.special_comfort_sent: + deferred.special_comfort_sent = True + self._record_inflight_backend_activity() + self._play_deferred_interruption_comfort() + log_flow_event( + self._call_logger, + "deferred_interruption_vad_promoted", + promoted_utterances=len(promoted), + promoted_long_stt_finals=long_count, + ) + + def _consume_vad_for_stt_final(self) -> DeferredVADUtterance | None: + deferred = self.state.deferred_interruption + self._prune_pending_vad_utterances() + queue = ( + deferred.pending_vad_utterances + if deferred.backend_in_flight or deferred.replay_in_flight + else deferred.pre_backend_vad_utterances + ) + if not queue: + return None + utterance = queue.popleft() + if deferred.backend_in_flight and utterance.is_long: + deferred.pending_long_stt_finals = max(0, deferred.pending_long_stt_finals - 1) + if not deferred.pending_long_stt_finals: + self._deferred_long_stt_settled.set() + return utterance + + def _reset_deferred_interruption(self) -> None: + deferred = self.state.deferred_interruption + deferred.backend_in_flight = False + deferred.pre_backend_vad_utterances.clear() + deferred.pending_vad_utterances.clear() + deferred.pending_long_stt_finals = 0 + deferred.long_turns.clear() + deferred.special_comfort_sent = False + deferred.replay_in_flight = False + deferred.cycle_seq += 1 + self._deferred_empty_final_latched_at = 0.0 + self._deferred_long_stt_settled.set() + + def _release_deferred_cycle(self, *, is_deferred_replay: bool) -> None: + """Fecha o ciclo diferido ao encerrar um turno do pipeline. + + R2 sempre limpa: e o fim do reenvio unico. R1 so limpa se chegou a abrir + o ciclo, e nunca quando acabou de despachar R2 -- ai a marca de reenvio + pertence ao turno seguinte, nao a este. + """ + deferred = self.state.deferred_interruption + if is_deferred_replay or ( + deferred.special_comfort_sent and not deferred.replay_in_flight + ): + self._reset_deferred_interruption() + + def _play_deferred_interruption_comfort(self, *, message_id: str = "") -> None: + deferred = self.state.deferred_interruption + cycle_seq = deferred.cycle_seq + + async def _play() -> None: + # O audio so entra no silencio do cliente. A espera tem teto porque + # um estado de fala travado deixaria a tarefa pendurada ate o fim da + # ligacao. + if not await self._wait_event_or_finalized( + self._user_not_speaking, + timeout=max( + 0.0, + float(self._config.deferred_interruption_user_turn_timeout_s), + ), + ): + log_flow_event( + self._call_logger, + "deferred_interruption_comfort_skipped", + reason="user_still_speaking", + ) + return + deferred = self.state.deferred_interruption + if ( + self.finalized.is_set() + or deferred.cycle_seq != cycle_seq + or deferred.replay_in_flight + or not deferred.special_comfort_sent + ): + return + base_message_id = str(message_id or "").strip() or self._current_base_message_id() + comfort_message_id = self._interruption_comfort_message_id( + message_id=base_message_id, + ) + audio_duration_ms = wav_duration_ms(str(DEFERRED_INTERRUPTION_COMFORT_AUDIO_PATH)) + await self.say_stage( + DEFERRED_INTERRUPTION_COMFORT_TEXT, + "INTERRUPTION_COMFORT", + # Comeca protegido contra o STT final da fala que originou este + # conforto e volta a ser interrompivel assim que esse final chega. + allow_interruptions=False, + rearm_interruptions_on_next_user_speech=True, + add_to_chat_ctx=False, + schedule_idle_after=False, + require_user_not_speaking=True, + message_id=comfort_message_id, + audio=wav_audio_frames(str(DEFERRED_INTERRUPTION_COMFORT_AUDIO_PATH)), + audio_duration_ms=audio_duration_ms, + ) + + self._create_task_logged( + _play(), + name="deferred_interruption_comfort", + ) + + def _play_deferred_short_comfort(self, *, source: str) -> None: + deferred = self.state.deferred_interruption + cycle_seq = deferred.cycle_seq + replay_in_flight = deferred.replay_in_flight + log_flow_event( + self._call_logger, + "deferred_short_comfort_scheduled", + source=source, + replay_in_flight=replay_in_flight, + ) + + async def _play() -> None: + current = self.state.deferred_interruption + if ( + self.finalized.is_set() + or current.cycle_seq != cycle_seq + or current.replay_in_flight != replay_in_flight + ): + return + await self._play_inflight_backend_wait_notice_audio( + user_seq=None, + attempt=1, + require_user_not_speaking=True, + ) + + self._create_task_logged( + _play(), + name=f"deferred_short_comfort_{source}", + ) + + async def _dispatch_deferred_interruption(self, reply: BackendReply) -> bool: + deferred = self.state.deferred_interruption + if deferred.replay_in_flight or not deferred.long_turns: + return False + turns = list(deferred.long_turns) + transcription, text = self._combined_interruption_transcription(turns) + turn = turns[-1] + speech_id = self._speech_id_from_reply(reply) + # O reenvio e unico: a marca precisa estar ativa antes de R2 ser + # agendado para que nenhum evento concorrente do cliente abra R3. + deferred.pending_vad_utterances.clear() + deferred.pending_long_stt_finals = 0 + deferred.long_turns.clear() + deferred.replay_in_flight = True + self._deferred_long_stt_settled.set() + await self.execute( + SetPendingInterrupt(listened_text="", skipped=True, speech_id=speech_id) + ) + log_flow_event( + self._call_logger, + "deferred_interruption_dispatched", + user_seq=turn.seq, + speech_id=speech_id, + accumulated_turns=len(turns), + text=text, + ) + self._create_task_logged( + self.run_pipeline( + transcription, + text, + is_deferred_replay=True, + user_seq=turn.seq, + message_id=turn.message_id, + inflight_initial_notice_consumed=True, + ), + name=f"run_pipeline_deferred_interruption_{turn.seq}", + ) + return True + + def note_vad_speech_end(self, speech_duration_ms: int) -> None: + duration_ms = max(0, int(speech_duration_ms or 0)) + self._last_vad_speech_end_at = time.monotonic() + self._last_vad_speech_duration_ms = duration_ms + self._emit_debug_event( + "user.speech.finished", + speech_duration_ms=duration_ms, + ) + deferred = self.state.deferred_interruption + minimum_audio_ms = max(0, int(self._config.deferred_interruption_min_audio_ms)) + is_long = duration_ms >= minimum_audio_ms + self._prune_pending_vad_utterances() + + if deferred.replay_in_flight: + log_flow_event( + self._call_logger, + "deferred_replay_vad_discarded", + speech_duration_ms=duration_ms, + is_long=is_long, + minimum_audio_ms=minimum_audio_ms, + short_comfort_scheduled=False, + ) + return + + utterance = DeferredVADUtterance( + duration_ms=duration_ms, + is_long=is_long, + ended_at=time.monotonic(), + ) + if not deferred.backend_in_flight: + if not self._config.deferred_interruption_enabled: + return + # O STT pode terminar depois que R1 for aberto. Guardamos o fim de + # fala para preservar a associacao VAD/STT nessa corrida. + deferred.pre_backend_vad_utterances.append(utterance) + log_flow_event( + self._call_logger, + "deferred_interruption_vad_pending_pre_backend", + speech_duration_ms=duration_ms, + is_long=is_long, + minimum_audio_ms=minimum_audio_ms, + ) + return + + if not self._config.deferred_interruption_enabled: + log_flow_event( + self._call_logger, + "backend_processing_vad_discarded", + speech_duration_ms=duration_ms, + minimum_audio_ms=minimum_audio_ms, + mode="feedback_compat", + ) + return + + deferred.pending_vad_utterances.append(utterance) + if is_long: + # Only utterances accepted by the business duration rule start a + # new interaction cycle. Short utterances may temporarily suppress + # playout while the customer is speaking, but must not postpone the + # periodic notice for the request already in flight. + self._record_inflight_backend_activity() + deferred.pending_long_stt_finals += 1 + self._deferred_long_stt_settled.clear() + if not deferred.special_comfort_sent: + deferred.special_comfort_sent = True + self._play_deferred_interruption_comfort() + else: + self._play_deferred_short_comfort(source="repeated_interruption") + + log_flow_event( + self._call_logger, + "deferred_interruption_vad_end", + speech_duration_ms=duration_ms, + is_long=is_long, + minimum_audio_ms=minimum_audio_ms, + ) + + def _close_deferred_vad_window(self) -> None: + """Descarta as falas VAD que nao receberam final dentro do ciclo. + + A associacao VAD/STT e por ordem de chegada. Uma fala sem final que fique + na fila seria consumida pelo final de outra fala no turno seguinte, e a + partir dai todo o pareamento anda deslocado. + """ + deferred = self.state.deferred_interruption + deferred.backend_in_flight = False + discarded_backend = len(deferred.pending_vad_utterances) + discarded_pre_backend = len(deferred.pre_backend_vad_utterances) + if discarded_backend or discarded_pre_backend: + log_flow_event( + self._call_logger, + "deferred_interruption_vad_window_closed", + discarded_utterances=discarded_backend + discarded_pre_backend, + discarded_backend_utterances=discarded_backend, + discarded_pre_backend_utterances=discarded_pre_backend, + pending_long_stt_finals=deferred.pending_long_stt_finals, + ) + deferred.pre_backend_vad_utterances.clear() + deferred.pending_vad_utterances.clear() + deferred.pending_long_stt_finals = 0 + self._deferred_long_stt_settled.set() + + async def _wait_event_or_finalized(self, event: asyncio.Event, *, timeout: float) -> bool: + """Espera o evento, mas desiste se a chamada terminar antes. + + Sem a corrida com `finalized` um desligamento no meio da espera ficaria + parado ate o teto do timeout. + """ + if event.is_set(): + return True + waiters = [ + asyncio.ensure_future(event.wait()), + asyncio.ensure_future(self.finalized.wait()), + ] + try: + await asyncio.wait( + waiters, timeout=timeout, return_when=asyncio.FIRST_COMPLETED + ) + finally: + for waiter in waiters: + if not waiter.done(): + waiter.cancel() + await asyncio.gather(*waiters, return_exceptions=True) + return event.is_set() + + async def _wait_deferred_interruption_settlement(self, *, user_seq: int) -> None: + """Segura a resposta de R1 ate o turno do cliente fechar. + + Com o VAD aberto nao ha o que decidir: a fala em curso ainda pode virar + mais uma transcricao do acumulado. Depois do END_OF_SPEECH resta esperar + o STT correspondente, e essa espera e curta e com teto -- a associacao + VAD/STT e por ordem de chegada, entao um final que nunca vem (falha do + provider, dois segmentos de VAD para um final so) prenderia o run_lock + pelo resto da ligacao. + """ + deferred = self.state.deferred_interruption + user_turn_timeout_s = max( + 0.0, float(self._config.deferred_interruption_user_turn_timeout_s) + ) + stt_settle_timeout_s = max( + 0.0, float(self._config.deferred_interruption_stt_settle_timeout_s) + ) + user_turn_deadline = time.monotonic() + user_turn_timeout_s + waited = False + + while not self.finalized.is_set(): + if not self._user_not_speaking.is_set(): + remaining = user_turn_deadline - time.monotonic() + if remaining <= 0.0: + log_flow_event( + self._call_logger, + "deferred_interruption_settle_timeout", + reason="user_turn", + user_seq=user_seq, + timeout_s=user_turn_timeout_s, + ) + break + waited = True + if not await self._wait_event_or_finalized( + self._user_not_speaking, timeout=remaining + ): + log_flow_event( + self._call_logger, + "deferred_interruption_settle_timeout", + reason="user_turn", + user_seq=user_seq, + timeout_s=user_turn_timeout_s, + ) + break + # O END_OF_SPEECH do VAD e a mudanca de estado do usuario nao tem + # ordem garantida; a folga evita liberar R1 no instante entre os + # dois e perder a ultima fala. + await asyncio.sleep(POST_USER_FINAL_SPEECH_GRACE_S) + continue + + if not deferred.pending_long_stt_finals: + break + + waited = True + if not await self._wait_event_or_finalized( + self._deferred_long_stt_settled, timeout=stt_settle_timeout_s + ): + log_flow_event( + self._call_logger, + "deferred_interruption_settle_timeout", + reason="stt_final", + user_seq=user_seq, + timeout_s=stt_settle_timeout_s, + pending_long_stt_finals=deferred.pending_long_stt_finals, + ) + # A fila esta dessincronizada: um final atrasado consumiria a + # fala errada mais adiante. + deferred.pending_vad_utterances.clear() + deferred.pending_long_stt_finals = 0 + self._deferred_long_stt_settled.set() + break + + if waited: + log_flow_event( + self._call_logger, + "deferred_interruption_settled", + user_seq=user_seq, + accumulated_turns=len(deferred.long_turns), + ) + + def _peek_started_turn_message_id(self) -> str: + return peek_started_turn_message_id(self._structured_log_context) + + def _clear_started_turn_message_id(self, message_id: str = "") -> None: + clear_started_turn_message_id( + self._structured_log_context, + message_id=message_id, + ) + + @staticmethod + def _base_message_id(message_id: str) -> str: + normalized_message_id = str(message_id or "").strip() + if not normalized_message_id: + return "" + + for marker in ( + "_conforto_", + "_inatividade_", + "_idle_nudge_", + "_interruption_confort_", + "_feedback_", + "_tts_error_", + "_transferencia_", + "_erro_terminal_", + "_erro_recurso_", + ): + base_message_id, separator, suffix = normalized_message_id.rpartition(marker) + if separator and base_message_id and suffix: + return base_message_id + return normalized_message_id + + + @staticmethod + def _feedback_message_id(*, message_id: str = "", speech_id: str = "") -> str: + base = str(message_id or "").strip() or str(uuid.uuid4()) + feedback_id = str(speech_id or "").strip() or uuid.uuid4().hex + return f"{base}_feedback_{feedback_id}" + + @staticmethod + def _idle_nudge_message_id(*, message_id: str = "", attempt: int) -> str: + base = str(message_id or "").strip() or str(uuid.uuid4()) + return f"{base}_idle_nudge_{max(1, int(attempt or 1))}" + + @staticmethod + def _interruption_comfort_message_id(*, message_id: str = "") -> str: + base = str(message_id or "").strip() or str(uuid.uuid4()) + return f"{base}_interruption_confort_{uuid.uuid4().hex}" + + @staticmethod + def _tts_error_message_id(*, message_id: str = "") -> str: + base = str(message_id or "").strip() or str(uuid.uuid4()) + return f"{base}_tts_error_{uuid.uuid4().hex}" + + @staticmethod + def _transfer_message_id(*, message_id: str = "") -> str: + base = str(message_id or "").strip() or str(uuid.uuid4()) + return f"{base}_transferencia_{uuid.uuid4().hex}" + + @staticmethod + def _resource_error_message_id( + *, + message_id: str = "", + terminal: bool, + tipo_evento: str, + ) -> str: + base = str(message_id or "").strip() or str(uuid.uuid4()) + family = "erro_terminal" if terminal else "erro_recurso" + direction = "recebimento" if tipo_evento == EVENT_RECEBIMENTO_MSG else "envio" + return f"{base}_{family}_{direction}_{uuid.uuid4().hex}" + + def _remember_agent_base_message_id(self, message_id: str) -> None: + base_message_id = self._base_message_id(message_id) + if base_message_id: + self._last_agent_message_id = base_message_id + + def _current_base_message_id(self) -> str: + return self._peek_started_turn_message_id() or self._last_agent_message_id + + def _resolve_agent_message_id(self, message_id: str = "", *, generate: bool = False) -> str: + resolved_message_id = str(message_id or "").strip() or self._current_base_message_id() + if not resolved_message_id and generate: + resolved_message_id = next_turn_message_id(self._structured_log_context) + self._remember_agent_base_message_id(resolved_message_id) + return resolved_message_id + + @classmethod + def _allow_interruptions_for_reply( + cls, + reply: BackendReply, + default: bool, + ) -> bool: + explicit = cls._metadata_bool(reply.metadata, "is_interruptible") + if explicit is None: + return default + return explicit + + def _current_speech_allows_interruption(self) -> bool: + if not self.speaking.is_set(): + return False + return bool(self.state.current_speech.allow_interruptions) + + def _rearm_current_speech_interruptions_on_new_user_speech(self) -> None: + current = self.state.current_speech + if ( + not self.speaking.is_set() + or not current.rearm_interruptions_on_next_user_speech + ): + return + + handle = current.handle + stage = current.stage + if handle is not None and bool(getattr(handle, "interrupted", False)): + return + try: + if handle is not None: + handle.allow_interruptions = True + except Exception: + self._logger.exception( + "[audio] failed rearming comfort interruptions | stage=%s", + stage, + ) + return + + current.rearm_interruptions_on_next_user_speech = False + current.allow_interruptions = True + log_flow_event( + self._call_logger, + "comfort_interruptions_rearmed", + stage=stage, + reason="new_user_speech", + ) + + @staticmethod + def _status_for_resource(resource: str) -> str: + normalized = str(resource or "").strip() + return IN_SESSION_STOP_STATUS_BY_RESOURCE.get(normalized, "stop_bridge_failed") + + @staticmethod + def _resource_from_session_error(error: Any) -> str: + error_type = str(getattr(error, "type", "") or "").strip().lower() + if error_type == "stt_error": + return "stt" + if error_type == "tts_error": + return "tts" + return "agent_runtime" + + @staticmethod + def _payload_requests_transfer(payload: Any) -> bool: + if isinstance(payload, Mapping): + for key in ("type", "status", "result_type", "resultType"): + value = str(payload.get(key) or "").strip().lower() + if value in {"transferred", "transferido"}: + return True + result = payload.get("result") + if result is not payload and CallRuntime._payload_requests_transfer(result): + return True + return False + + if isinstance(payload, (list, tuple)): + return any(CallRuntime._payload_requests_transfer(item) for item in payload) + + return False + + @classmethod + def _error_message_from_backend_reply(cls, reply: BackendReply) -> str | None: + if cls._payload_requests_transfer(reply.export_payload): + return error_message_from_resource( + status="transferred", + tipo_evento=EVENT_ENVIO_MSG, + ) + return None + + def _log_backend_reply_error_event( + self, + reply: BackendReply, + *, + source: str, + message_id: str = "", + ) -> None: + transfer_status = "" + if isinstance(reply.export_payload, Mapping): + transfer_status = str(reply.export_payload.get("status") or "").strip() + + self._log_structured_error_event( + tipo_evento=EVENT_ENVIO_MSG, + erro_msg=self._error_message_from_backend_reply(reply), + message_id=self._transfer_message_id( + message_id=message_id or self._message_id_from_reply(reply) + ), + status=transfer_status, + reason="agent_transfer", + source=source, + ) + + def _log_structured_error_event( + self, + *, + tipo_evento: str, + erro_msg: str | None, + erro_detalhe: str | None = None, + message_id: str = "", + status: str = "", + reason: str = "", + resource: str = "", + source: str = "", + http_cod_status: int | str | None = None, + http_cod_desc: str | None = None, + ) -> None: + erro_msg = str(erro_msg or "").strip() + if not erro_msg: + return + + dedupe_key = ( + str(tipo_evento or ""), + erro_msg, + str(status or ""), + str(reason or ""), + str(resource or ""), + str(message_id or ""), + ) + if dedupe_key in self._structured_error_event_keys: + return + self._structured_error_event_keys.add(dedupe_key) + + now_ns = time.time_ns() + log_structured_event( + self._call_logger, + self._structured_log_context, + tipo_evento=tipo_evento, + message_id=message_id, + inicio_ns=now_ns, + fim_ns=now_ns, + latencia_total_ms=0, + latencia_tffb_ms=0, + erro_msg=erro_msg, + erro_detalhe=erro_detalhe, + http_cod_status=http_cod_status, + http_cod_desc=http_cod_desc, + finalizacao=status or None, + ) + log_flow_event( + self._call_logger, + "structured_error", + tipo_evento=tipo_evento, + erro_msg=erro_msg, + status=status, + reason=reason, + resource=resource, + source=source, + ) + + def _log_resource_error_events( + self, + *, + status: str = "", + reason: str = "", + resource: str = "", + source: str = "", + message_id: str = "", + exc: Exception | None = None, + terminal: bool = False, + ) -> None: + message_id = str(message_id or "").strip() or self._current_base_message_id() + detail = f"{type(exc).__name__}: {exc}" if exc is not None else None + for tipo_evento in (EVENT_RECEBIMENTO_MSG, EVENT_ENVIO_MSG): + erro_msg = error_message_from_resource( + resource=resource, + status=status, + reason=reason, + tipo_evento=tipo_evento, + ) + self._log_structured_error_event( + tipo_evento=tipo_evento, + erro_msg=erro_msg, + erro_detalhe=detail, + status=status, + reason=reason, + resource=resource, + source=source, + message_id=self._resource_error_message_id( + message_id=message_id, + terminal=terminal, + tipo_evento=tipo_evento, + ), + ) + + async def notify_bridge_stop( + self, + *, + status: str, + reason: str, + resource: str = "", + failed_resources: tuple[str, ...] = (), + phase: str = "in_session", + ) -> None: + try: + await self.execute( + NotifyBridgeStop( + status=status, + reason=reason, + resource=resource, + failed_resources=failed_resources, + phase=phase, + ) + ) + if self._timeline is not None: + self._timeline.emit( + "bridge_stop_requested", + status=status, + reason=reason, + resource=resource, + failed_resources=list(failed_resources), + phase=phase, + ) + except Exception as exc: + self._log_resource_error_events( + status="stop_bridge_failed", + reason="bridge_notify_stop_failed", + resource="bridge", + source="notify_bridge_stop", + exc=exc, + ) + self._logger.exception("[bridge] Falha enviando STOP") + + async def terminate_with_resource_stop( + self, + *, + resource: str, + source: str, + reason: str = "resource_unhealthy", + ) -> None: + async with self._terminal_stop_lock: + if self._terminal_stop_sent: + return + self._terminal_stop_sent = True + + status = self._status_for_resource(resource) + if self._timeline is not None: + self._timeline.emit( + "terminal_stop_started", + status=status, + reason=reason, + resource=resource, + source=source, + ) + + await self.notify_bridge_stop( + status=status, + reason=reason, + resource=resource, + failed_resources=(resource,) if resource else (), + phase="in_session", + ) + self._log_resource_error_events( + status=status, + reason=reason, + resource=resource, + source=source, + terminal=True, + ) + await self.finalize(status) + + async def _wait_post_user_final_grace(self, *, stage: str) -> None: + last_user_final_at = float(self.state.last_user_final_at or 0.0) + if last_user_final_at <= 0.0: + return + + remaining_s = POST_USER_FINAL_SPEECH_GRACE_S - (time.monotonic() - last_user_final_at) + if remaining_s <= 0.0: + return + + self._call_logger.info( + "BOT_SAY_WAIT_POST_FINAL | stage=%s | wait_ms=%s", + stage, + round(remaining_s * 1000), + ) + if self._timeline is not None: + self._timeline.emit( + "tts_stage_wait_post_final", + stage=stage, + wait_ms=round(remaining_s * 1000), + ) + await asyncio.sleep(remaining_s) + + def _has_remote_participants(self) -> bool: + participants = getattr(getattr(self._ctx, "room", None), "remote_participants", None) + return bool(participants) + + @staticmethod + def _decode_data_packet(data_packet: Any) -> dict[str, Any] | None: + raw = getattr(data_packet, "data", "") + if isinstance(raw, (bytes, bytearray)): + raw = raw.decode("utf-8", "ignore") + if not isinstance(raw, str): + return None + + message = raw.strip() + if not message: + return None + + try: + payload = json.loads(message) + except Exception: + return None + + if not isinstance(payload, dict): + return None + return payload + + @classmethod + def _user_audio_input_setup_timeout_s(cls) -> float: + return cls._float_env_s( + "USER_AUDIO_INPUT_SETUP_TIMEOUT_S", + USER_AUDIO_INPUT_SETUP_TIMEOUT_S, + ) + + def _set_session_audio_input_enabled(self, enabled: bool) -> bool: + session_input = getattr(self._session, "input", None) + set_audio_enabled = getattr(session_input, "set_audio_enabled", None) + if not callable(set_audio_enabled): + return False + + try: + set_audio_enabled(enabled) + except Exception: + self._logger.exception("[audio] failed toggling session audio input") + return False + return True + + def hold_user_audio_input(self, *, reason: str) -> None: + """Descarta o audio do usuario na origem ate a primeira fala do agente terminar.""" + if self._user_audio_input_gate_released or not self._user_audio_input_enabled: + return + if not self._set_session_audio_input_enabled(False): + return + + self._user_audio_input_enabled = False + self._call_logger.info("USER_AUDIO_INPUT_HELD | reason=%s", reason) + log_flow_event( + self._call_logger, + "user_audio_input_held", + reason=reason, + stage=self.state.current_stage, + ) + if self._timeline is not None: + self._timeline.emit("user_audio_input_held", reason=reason) + + def release_user_audio_input(self, *, reason: str) -> None: + if self._user_audio_input_gate_released: + return + + self._user_audio_input_gate_released = True + self._scheduler.cancel( + "user_audio_input_gate", + reason=reason, + log_name="USER_AUDIO_INPUT_GATE_CANCEL", + ) + if self._user_audio_input_enabled: + return + + self._set_session_audio_input_enabled(True) + self._user_audio_input_enabled = True + self._call_logger.info("USER_AUDIO_INPUT_RELEASED | reason=%s", reason) + log_flow_event( + self._call_logger, + "user_audio_input_released", + reason=reason, + stage=self.state.current_stage, + ) + if self._timeline is not None: + self._timeline.emit("user_audio_input_released", reason=reason) + + def _arm_user_audio_input_gate_timeout(self) -> None: + """Rede de seguranca: sem o turno inicial, o gate nunca seria liberado.""" + delay_s = self._user_audio_input_setup_timeout_s() + if delay_s <= 0: + return + + token = 0 + + async def _timer() -> None: + try: + await asyncio.sleep(delay_s) + except asyncio.CancelledError: + return + + if not self._scheduler.is_current("user_audio_input_gate", token): + return + if self._initial_agent_turn_requested: + return + + self._call_logger.warning( + "USER_AUDIO_INPUT_GATE_TIMEOUT | timeout_s=%.1f | reason=initial_agent_turn_not_requested", + delay_s, + ) + self.release_user_audio_input(reason="setup_timeout") + + token = self._scheduler.arm( + "user_audio_input_gate", + task_name="user_audio_input_gate_timeout", + coro=_timer(), + ) + + def request_initial_agent_turn(self, *, reason: str) -> None: + if not self._agent_starts_conversation: + return + if self._initial_agent_turn_requested or self.finalized.is_set(): + return + if self.state.user_final_seq != 0: + return + if str(self.state.current_stage or "INTRO").upper() not in {"", "INTRO"}: + return + + self._initial_agent_turn_requested = True + self._scheduler.cancel( + "user_audio_input_gate", + reason="initial_agent_turn_requested", + log_name="USER_AUDIO_INPUT_GATE_CANCEL", + ) + self._call_logger.info("INITIAL_AGENT_TURN | requested | reason=%s", reason) + if self._timeline is not None: + self._timeline.emit("initial_agent_turn_requested", reason=reason) + + async def _run_initial_turn() -> None: + try: + await self.run_pipeline("", "", user_seq=0) + except Exception: + self._logger.exception("[pipeline] initial agent turn failed") + if self._timeline is not None: + self._timeline.emit("initial_agent_turn_failed", reason=reason) + finally: + self.release_user_audio_input(reason="first_agent_message_played") + + self._create_task_logged( + _run_initial_turn(), + name="run_pipeline_initial", + ) + + def _supports_backend_push(self) -> bool: + pipeline = self._agent.pipeline + if pipeline is None: + return False + + supports = getattr(pipeline, "supports_server_push", None) + if callable(supports): + try: + return bool(supports()) + except Exception: + return False + + return callable(getattr(pipeline, "wait_for_server_push", None)) + + def _supports_inflight_backend_push(self) -> bool: + pipeline = self._agent.pipeline + if pipeline is None: + return False + + supports = getattr(pipeline, "supports_inflight_backend_push", None) + if not callable(supports): + return False + + try: + return bool(supports()) + except Exception: + return False + + def _inflight_backend_wait_enabled(self) -> bool: + return ( + self._supports_inflight_backend_push() + and float(self._config.inflight_backend_wait_timeout_s or 0.0) > 0.0 + ) + + def _inflight_backend_wait_notice_enabled(self) -> bool: + return ( + self._inflight_backend_wait_enabled() + and self._has_inflight_backend_wait_audio() + and float(self._config.inflight_backend_wait_interval_s or 0.0) > 0.0 + ) + + def _record_inflight_backend_activity(self) -> None: + self._inflight_backend_activity_at = time.monotonic() + + def _should_continue_inflight_backend_wait(self) -> bool: + """Keep supervising the backend task for its whole lifetime. + + A valid deferred interruption increments ``user_final_seq`` while the + original backend request remains in flight. Tying this supervisor to + that sequence used to disable both periodic notices and the timeout, + leaving the call waiting silently for the backend. The coroutine that + owns the backend task already defines the request lifetime; only call + finalization should stop its supervision early. + """ + return not self.finalized.is_set() + + @staticmethod + def _resolve_wait_audio_path(raw_path: str) -> Path: + path = Path(raw_path).expanduser() + if not path.is_absolute(): + path = Path.cwd() / path + return path + + @classmethod + def _wait_audio_candidates_from_value(cls, value: str) -> tuple[Path, ...]: + raw_path = str(value or "").strip() + if not raw_path: + return () + + path = cls._resolve_wait_audio_path(raw_path) + if path.is_file() and path.suffix.lower() == ".wav": + return (path,) + if not path.is_dir(): + return () + + return tuple( + sorted( + child + for child in path.iterdir() + if child.is_file() and child.suffix.lower() == ".wav" + ) + ) + + @classmethod + def _inflight_backend_wait_audio_candidates_from_config( + cls, + config: RuntimeConfig, + ) -> dict[str, tuple[Path, ...]]: + long_candidates: list[Path] = [] + long_candidates.extend( + cls._wait_audio_candidates_from_value( + config.inflight_backend_wait_long_audio_dir + ) + ) + long_candidates.extend( + cls._wait_audio_candidates_from_value( + config.inflight_backend_wait_long_audio_path + ) + ) + return { + "short": cls._wait_audio_candidates_from_value( + config.inflight_backend_wait_short_audio_dir + ), + "long": tuple(dict.fromkeys(long_candidates)), + } + + @staticmethod + def _unique_wait_audio_candidates( + candidates_by_kind: Mapping[str, Iterable[Path]], + ) -> tuple[Path, ...]: + paths: list[Path] = [] + for candidates in candidates_by_kind.values(): + paths.extend(candidates) + return tuple(dict.fromkeys(paths)) + + @staticmethod + def _safe_log_warning(logger: Any, msg: str, *args: Any) -> None: + log_fn = getattr(logger, "warning", None) or getattr(logger, "info", None) + if not callable(log_fn): + return + try: + log_fn(msg, *args) + except Exception: + return + + @classmethod + def _preload_wait_audio_paths( + cls, + paths: Iterable[Path], + *, + logger: Any = None, + ) -> tuple[dict[Path, int], int]: + durations: dict[Path, int] = {} + failures = 0 + for audio_path in paths: + try: + durations[audio_path] = wav_duration_ms(str(audio_path)) + except Exception as exc: + failures += 1 + cls._safe_log_warning( + logger, + "[audio] failed preloading wait audio path=%s error=%s: %s", + audio_path, + type(exc).__name__, + exc, + ) + return durations, failures + + @classmethod + def prewarm_wait_audio_cache( + cls, + *, + short_audio_dir: str = "", + long_audio_dir: str = "", + long_audio_path: str = "", + logger: Any = None, + ) -> tuple[int, int]: + config = RuntimeConfig( + call_end_grace_s=0.0, + final_grace_s=0.0, + idle_nudge_delay_s=0.0, + idle_nudge_join_delay_s=0.0, + idle_nudge_close_delay_s=0.0, + idle_nudge_max_tries=0, + idle_nudge_end_reason="", + inflight_backend_wait_short_audio_dir=short_audio_dir, + inflight_backend_wait_long_audio_dir=long_audio_dir, + inflight_backend_wait_long_audio_path=long_audio_path, + ) + candidates_by_kind = cls._inflight_backend_wait_audio_candidates_from_config(config) + durations, failures = cls._preload_wait_audio_paths( + cls._unique_wait_audio_candidates(candidates_by_kind), + logger=logger, + ) + return len(durations), failures + + def _inflight_backend_wait_audio_candidates(self, kind: str) -> tuple[Path, ...]: + normalized = str(kind or "").strip().lower() + if normalized == "short": + return self._inflight_backend_wait_audio_candidates_cache.get("short", ()) + if normalized == "long": + return self._inflight_backend_wait_audio_candidates_cache.get("long", ()) + return () + + def _has_inflight_backend_wait_audio(self) -> bool: + return bool( + self._inflight_backend_wait_audio_candidates("short") + or self._inflight_backend_wait_audio_candidates("long") + ) + + @staticmethod + def _preferred_inflight_backend_wait_audio_kind(attempt: int) -> str: + return "short" if int(attempt or 0) <= 1 else "long" + + def _select_inflight_backend_wait_audio( + self, + *, + attempt: int, + ) -> tuple[Path | None, str]: + preferred_kind = self._preferred_inflight_backend_wait_audio_kind(attempt) + candidates = self._inflight_backend_wait_audio_candidates(preferred_kind) + selected_kind = preferred_kind + if not candidates: + selected_kind = "long" if preferred_kind == "short" else "short" + candidates = self._inflight_backend_wait_audio_candidates(selected_kind) + + if not candidates: + return None, selected_kind + return random.choice(candidates), selected_kind + + def _inflight_backend_wait_audio_path(self, *, attempt: int = 2) -> Path | None: + audio_path, _kind = self._select_inflight_backend_wait_audio(attempt=attempt) + return audio_path + + def _default_inflight_backend_wait_text(self) -> str: + text = str(self._config.inflight_backend_wait_text or "").strip() + return text or "Um momento, ainda estou consultando para te ajudar." + + @staticmethod + def _inflight_backend_wait_audio_text(audio_path: Path, *, fallback_text: str) -> str: + text_path = audio_path.with_suffix(".txt") + if not text_path.is_file(): + return fallback_text + + for encoding in ("utf-8-sig", "utf-8", "latin-1"): + try: + text = " ".join(text_path.read_text(encoding=encoding).split()) + except UnicodeDecodeError: + continue + except OSError: + return fallback_text + return text or fallback_text + + return fallback_text + + def _cached_inflight_backend_wait_audio_text( + self, + audio_path: Path, + *, + fallback_text: str, + ) -> str: + key = (audio_path, fallback_text) + cached = self._inflight_backend_wait_audio_text_cache.get(key) + if cached is not None: + return cached + text = self._inflight_backend_wait_audio_text( + audio_path, + fallback_text=fallback_text, + ) + self._inflight_backend_wait_audio_text_cache[key] = text + return text + + def _preload_inflight_backend_wait_audio_texts(self) -> None: + fallback_text = self._default_inflight_backend_wait_text() + for audio_path in self._unique_wait_audio_candidates( + self._inflight_backend_wait_audio_candidates_cache + ): + self._cached_inflight_backend_wait_audio_text( + audio_path, + fallback_text=fallback_text, + ) + + def _inflight_backend_wait_audio_duration_ms(self, audio_path: Path) -> int | None: + if audio_path in self._inflight_backend_wait_audio_duration_cache: + return self._inflight_backend_wait_audio_duration_cache[audio_path] + try: + duration_ms = wav_duration_ms(str(audio_path)) + except Exception: + return None + self._inflight_backend_wait_audio_duration_cache[audio_path] = duration_ms + return duration_ms + + @staticmethod + def _inflight_backend_wait_message_id(*, message_id: str = "", attempt: int) -> str: + base = str(message_id or "").strip() or str(uuid.uuid4()) + comfort_id = uuid.uuid4().hex + return f"{base}_conforto_{comfort_id}" + + @staticmethod + def _agent_wait_timeout_inactivity_message_id(*, message_id: str = "", attempt: int) -> str: + base = str(message_id or "").strip() or str(uuid.uuid4()) + return f"{base}_inatividade_{max(1, int(attempt or 1))}" + + async def _play_inflight_backend_wait_notice_audio( + self, + *, + user_seq: int | None, + attempt: int, + message_id: str = "", + require_user_not_speaking: bool = True, + abort_if_agent_spoke_since: int | None = None, + ) -> InflightBackendWaitNoticeResult: + if user_seq is not None and not self._should_continue_inflight_backend_wait(): + return InflightBackendWaitNoticeResult(started=False) + + processing_notice = user_seq is not None + fallback_text = self._default_inflight_backend_wait_text() + audio_path, notice_kind = self._select_inflight_backend_wait_audio(attempt=attempt) + if audio_path is None or not self._has_remote_participants(): + return InflightBackendWaitNoticeResult(started=False) + if require_user_not_speaking and not self._user_not_speaking.is_set(): + return InflightBackendWaitNoticeResult( + started=False, + interrupted_by_user=processing_notice, + ) + text = self._cached_inflight_backend_wait_audio_text( + audio_path, + fallback_text=fallback_text, + ) + duration_ms = self._inflight_backend_wait_audio_duration_ms(audio_path) + + base_message_id = str(message_id or "").strip() or self._current_base_message_id() + comfort_message_id = self._inflight_backend_wait_message_id( + message_id=base_message_id, + attempt=attempt, + ) + self._call_logger.info( + "INFLIGHT_BACKEND_WAIT_AUDIO | message_id=%s | original_message_id=%s | user_seq=%s | attempt=%s | kind=%s | path=%s | text=%r", + comfort_message_id, + base_message_id, + user_seq, + attempt, + notice_kind, + audio_path, + text, + ) + if self._timeline is not None: + self._timeline.emit( + "inflight_backend_wait_audio", + message_id=comfort_message_id, + user_seq=user_seq, + attempt=attempt, + kind=notice_kind, + path=str(audio_path), + duration_ms=duration_ms, + text=text, + ) + + interrupted_by_user = False + + def _capture_playout_result( + *, + interrupted: bool, + user_spoke_during: bool, + ) -> None: + nonlocal interrupted_by_user + interrupted_by_user = processing_notice and ( + interrupted or user_spoke_during + ) + + try: + started = await self.say_stage( + text, + INFLIGHT_BACKEND_WAIT_STAGE, + add_to_chat_ctx=False, + # O aviso curto comeca antes do STT final do mesmo turno. Ele fica + # protegido somente ate esse final chegar; depois uma nova fala + # real do cliente pode interrompe-lo normalmente. + allow_interruptions=notice_kind != "short", + rearm_interruptions_on_next_user_speech=notice_kind == "short", + schedule_idle_after=False, + wait_post_final_grace=False, + require_user_not_speaking=require_user_not_speaking, + abort_if_agent_spoke_since=abort_if_agent_spoke_since, + message_id=comfort_message_id, + audio=wav_audio_frames(str(audio_path)), + audio_duration_ms=duration_ms, + playout_result_callback=_capture_playout_result, + ) + result = InflightBackendWaitNoticeResult( + started=bool(started), + interrupted_by_user=interrupted_by_user, + ) + if result: + self._record_inflight_backend_activity() + return result + except Exception: + self._logger.exception("[audio] failed playing inflight backend wait audio") + if self._timeline is not None: + self._timeline.emit( + "audio_stage_failed", + stage=INFLIGHT_BACKEND_WAIT_STAGE, + source="inflight_backend_wait_audio", + path=str(audio_path), + ) + await self.terminate_with_resource_stop( + resource="tts", + source="inflight_backend_wait_audio", + ) + raise InflightBackendWaitTimedOut() + + def _cancel_pre_backend_wait_notice(self, *, reason: str) -> None: + self._pre_backend_wait_notice_generation += 1 + self._pre_backend_wait_notice_reserved = False + task = self._pre_backend_wait_notice_task + self._pre_backend_wait_notice_task = None + if task is not None and task is not asyncio.current_task() and not task.done(): + task.cancel() + if self._timeline is not None: + self._timeline.emit("pre_backend_wait_notice_cancelled", reason=reason) + + def _reserve_pre_backend_wait_notice(self) -> bool: + if ( + self.finalized.is_set() + or not self._inflight_backend_wait_notice_enabled() + or not self._has_remote_participants() + ): + return False + + self._pre_backend_wait_notice_reserved = True + self._record_inflight_backend_activity() + return True + + def _consume_pre_backend_wait_notice(self) -> bool: + if not self._pre_backend_wait_notice_reserved: + return False + + self._pre_backend_wait_notice_reserved = False + return True + + def _pre_backend_wait_notice_is_current_speech(self) -> bool: + return ( + self._pre_backend_wait_notice_reserved + and self.state.current_speech.stage == INFLIGHT_BACKEND_WAIT_STAGE + ) + + def _inflight_backend_wait_notice_is_speaking(self) -> bool: + return ( + self.speaking.is_set() + and self.state.current_speech.stage == INFLIGHT_BACKEND_WAIT_STAGE + ) + + def _pre_backend_wait_notice_fast_on_vad_pause(self, *, reason: str) -> bool: + return ( + bool(self._config.pre_backend_wait_notice_fast_on_vad_pause) + and str(reason or "").strip() == "vad_end_of_speech" + ) + + def _log_pre_backend_wait_notice_skipped(self, *, reason: str, skipped_reason: str) -> None: + log_flow_event( + self._call_logger, + "pre_backend_wait_notice_skipped", + reason=reason, + skipped_reason=skipped_reason, + ) + if self._timeline is not None: + self._timeline.emit( + "pre_backend_wait_notice_skipped", + reason=reason, + skipped_reason=skipped_reason, + ) + + def _schedule_pre_backend_wait_notice(self, *, reason: str) -> None: + if not self._inflight_backend_wait_notice_enabled() or self.finalized.is_set(): + return + post_playout_drop_context = self._active_post_playout_drop_context( + reason="protected_speech_post_playout_grace" + ) + if ( + self._drop_next_user_final_context is not None + or self._backend_processing_drop_context is not None + or post_playout_drop_context is not None + ): + self._log_pre_backend_wait_notice_skipped( + reason=reason, + skipped_reason="user_input_drop_pending", + ) + return + if self.speaking.is_set(): + # Nao ha silencio para mascarar: o aviso so entraria na fila do + # say_lock e sairia colado no fim da fala do agente. + self._log_pre_backend_wait_notice_skipped( + reason=reason, + skipped_reason="agent_speaking", + ) + return + if self._agent._run_lock.locked(): + # O fim de fala aconteceu enquanto o backend ainda processava o + # turno anterior. O STT pode transformar essa fala em uma + # interrupcao diferida, que possui seu proprio conforto dedicado. + # Tocar o aviso especulativo agora produziria dois confortos em + # sequencia: AGENT_BACKEND_WAIT e, logo depois, + # INTERRUPTION_COMFORT. + self._log_pre_backend_wait_notice_skipped( + reason=reason, + skipped_reason="backend_processing", + ) + return + if self._pre_backend_wait_notice_reserved: + return + if ( + self._pre_backend_wait_notice_task is not None + and not self._pre_backend_wait_notice_task.done() + ): + return + + self._cancel_pre_backend_wait_notice(reason=f"restart:{reason}") + self._pre_backend_wait_notice_generation += 1 + generation = self._pre_backend_wait_notice_generation + fast_on_vad_pause = self._pre_backend_wait_notice_fast_on_vad_pause( + reason=reason + ) + say_stage_seq_at_schedule = self._say_stage_seq + + async def _run() -> None: + try: + if not fast_on_vad_pause: + await asyncio.sleep(PRE_BACKEND_WAIT_NOTICE_GUARD_S) + if generation != self._pre_backend_wait_notice_generation: + return + if not fast_on_vad_pause and not self._user_not_speaking.is_set(): + self._log_pre_backend_wait_notice_skipped( + reason=reason, + skipped_reason="user_speaking", + ) + return + if self.speaking.is_set(): + self._log_pre_backend_wait_notice_skipped( + reason=reason, + skipped_reason="agent_speaking", + ) + return + if not self._reserve_pre_backend_wait_notice(): + return + + if self._timeline is not None: + self._timeline.emit( + "pre_backend_wait_notice_started", + reason=reason, + fast_on_vad_pause=fast_on_vad_pause, + ) + + played = await self._play_inflight_backend_wait_notice_audio( + user_seq=None, + attempt=1, + message_id=self._peek_started_turn_message_id(), + require_user_not_speaking=not fast_on_vad_pause, + abort_if_agent_spoke_since=say_stage_seq_at_schedule, + ) + if not played: + # Sem audio reproduzido nao ha aviso a consumir no pipeline: + # devolve a reserva para nao bloquear o proximo agendamento. + self._cancel_pre_backend_wait_notice(reason="notice_not_played") + except asyncio.CancelledError: + return + except InflightBackendWaitTimedOut: + return + finally: + if self._pre_backend_wait_notice_task is asyncio.current_task(): + self._pre_backend_wait_notice_task = None + + self._pre_backend_wait_notice_task = self._create_task_logged( + _run(), + name="pre_backend_wait_notice", + ) + + def schedule_pre_backend_wait_notice(self, *, reason: str) -> None: + self._schedule_pre_backend_wait_notice(reason=reason) + + async def _execute_pipeline_with_inflight_backend_wait( + self, + stt_payload: Any, + *, + user_seq: int, + message_id: str = "", + initial_notice_consumed: bool = False, + ) -> BackendReply: + if not self._inflight_backend_wait_enabled(): + return await self.execute(RunPipelineInput(stt_payload)) + + timeout_s = max(0.0, float(self._config.inflight_backend_wait_timeout_s or 0.0)) + interval_s = max(0.0, float(self._config.inflight_backend_wait_interval_s or 0.0)) + max_notices = max(0, int(self._config.inflight_backend_wait_max_notices or 0)) + can_notice = self._inflight_backend_wait_notice_enabled() + + def _notice_limit_reached(notices: int) -> bool: + # A non-positive limit explicitly means unlimited notices. The + # absolute backend timeout remains the terminal boundary. + return max_notices > 0 and notices >= max_notices + + started_at = time.monotonic() + backend_task = self._create_task_logged( + self.execute(RunPipelineInput(stt_payload)), + name=f"run_pipeline_backend_{user_seq}", + ) + notices_sent = 1 if initial_notice_consumed else 0 + pre_notice_consumed = initial_notice_consumed + max_notices_logged = False + user_silence_task: asyncio.Task[bool] | None = None + if initial_notice_consumed: + # O INTERRUPTION_COMFORT ja confirmou a fala do usuario. Conta como + # primeiro aviso para que o pipeline substituto nao toque um short + # logo em seguida; o proximo aviso, se necessario, sera long e + # respeitara o intervalo normal. + self._inflight_backend_activity_at = started_at + + def _notice_activity_at() -> float: + if pre_notice_consumed: + return self._inflight_backend_activity_at or started_at + return max(started_at, self._inflight_backend_activity_at) + + def _defer_notice_until_user_silence() -> None: + nonlocal user_silence_task + # Any real user activity, including a short utterance that business + # rules later discard, must rearm the notice interval. Otherwise an + # already-due notice is retried immediately after every silence and + # several comfort audios appear back-to-back. + self._record_inflight_backend_activity() + if user_silence_task is None or user_silence_task.done(): + user_silence_task = asyncio.create_task( + self._user_not_speaking.wait(), + name="inflight_backend_wait_user_silence", + ) + + def _log_max_notices_reached() -> None: + nonlocal max_notices_logged + if max_notices_logged or not can_notice or not _notice_limit_reached(notices_sent): + return + max_notices_logged = True + elapsed_s = max(0.0, time.monotonic() - started_at) + remaining_timeout_s = max(0.0, timeout_s - elapsed_s) + log_flow_event( + self._call_logger, + "inflight_backend_wait_max_notices_reached", + user_seq=user_seq, + max_notices=max_notices, + notices_sent=notices_sent, + wait_timeout_s=timeout_s, + remaining_timeout_s=round(remaining_timeout_s, 3), + ) + if self._timeline is not None: + self._timeline.emit( + "inflight_backend_wait_max_notices_reached", + user_seq=user_seq, + max_notices=max_notices, + notices_sent=notices_sent, + wait_timeout_s=timeout_s, + remaining_timeout_s=round(remaining_timeout_s, 3), + ) + self._log_resource_error_events( + status=self._status_for_resource("agent_backend"), + reason="inflight_backend_wait_max_notices_reached", + resource="agent_backend", + source="inflight_backend_wait_max_notices_reached", + message_id=message_id, + ) + + try: + if can_notice: + if initial_notice_consumed: + pass + elif self._consume_pre_backend_wait_notice(): + notices_sent = 1 + pre_notice_consumed = True + else: + self._inflight_backend_activity_at = started_at + await asyncio.sleep(0) + if backend_task.done(): + return backend_task.result() + + notices_sent += 1 + notice_result = await self._play_inflight_backend_wait_notice_audio( + user_seq=user_seq, + attempt=notices_sent, + message_id=message_id, + ) + if not notice_result: + notices_sent -= 1 + if notice_result.interrupted_by_user: + _defer_notice_until_user_silence() + else: + self._record_inflight_backend_activity() + else: + _log_max_notices_reached() + else: + self._inflight_backend_activity_at = started_at + + _log_max_notices_reached() + + while True: + if not self._should_continue_inflight_backend_wait(): + return await backend_task + + now = time.monotonic() + timeout_at = started_at + timeout_s + wait_until = timeout_at + pre_notice_task = self._pre_backend_wait_notice_task + pre_notice_running = ( + notices_sent > 0 + and pre_notice_task is not None + and not pre_notice_task.done() + ) + waiting_for_user_silence = ( + user_silence_task is not None + and not user_silence_task.done() + ) + if ( + can_notice + and not _notice_limit_reached(notices_sent) + and not pre_notice_running + and not waiting_for_user_silence + ): + last_activity_at = _notice_activity_at() + wait_until = min(wait_until, last_activity_at + interval_s) + + wait_s = max(0.0, wait_until - now) + wait_tasks = {backend_task} + if pre_notice_running: + wait_tasks.add(pre_notice_task) + if waiting_for_user_silence and user_silence_task is not None: + wait_tasks.add(user_silence_task) + + done, _pending = await asyncio.wait( + wait_tasks, + timeout=wait_s, + return_when=asyncio.FIRST_COMPLETED, + ) + if backend_task in done: + return backend_task.result() + if pre_notice_task is not None and pre_notice_task in done: + await asyncio.gather(pre_notice_task, return_exceptions=True) + continue + if user_silence_task is not None and user_silence_task in done: + await asyncio.gather(user_silence_task, return_exceptions=True) + user_silence_task = None + continue + + now = time.monotonic() + if timeout_s > 0.0 and now >= timeout_at: + if self._timeline is not None: + self._timeline.emit( + "inflight_backend_wait_timeout", + user_seq=user_seq, + wait_timeout_s=timeout_s, + notices_sent=notices_sent, + ) + backend_task.cancel() + await asyncio.gather(backend_task, return_exceptions=True) + await self.terminate_with_resource_stop( + resource="agent_backend", + source="inflight_backend_wait_timeout", + ) + raise InflightBackendWaitTimedOut() + + if can_notice and not _notice_limit_reached(notices_sent): + last_activity_at = _notice_activity_at() + if now - last_activity_at >= interval_s: + notices_sent += 1 + notice_result = await self._play_inflight_backend_wait_notice_audio( + user_seq=user_seq, + attempt=notices_sent, + message_id=message_id, + ) + if not notice_result: + notices_sent -= 1 + if notice_result.interrupted_by_user: + _defer_notice_until_user_silence() + else: + self._record_inflight_backend_activity() + else: + _log_max_notices_reached() + except InflightBackendWaitTimedOut: + if not backend_task.done(): + backend_task.cancel() + await asyncio.gather(backend_task, return_exceptions=True) + raise + except asyncio.CancelledError: + backend_task.cancel() + await asyncio.gather(backend_task, return_exceptions=True) + raise + finally: + if user_silence_task is not None and not user_silence_task.done(): + user_silence_task.cancel() + await asyncio.gather(user_silence_task, return_exceptions=True) + + def _cancel_backend_push_listener(self, *, reason: str, force: bool = False) -> None: + task = self._backend_push_task + if not force and self.speaking.is_set(): + return + self._backend_push_task = None + if task is not None and not task.done(): + task.cancel() + if self._timeline is not None: + self._timeline.emit("backend_push_cancelled", reason=reason) + + def _schedule_backend_push_listener( + self, + *, + user_seq: int, + reason: str, + inflight: bool = False, + ) -> None: + if self.finalized.is_set(): + return + + supported = ( + self._supports_inflight_backend_push() + if inflight + else self._supports_backend_push() + ) + if not supported: + if not inflight: + self._cancel_backend_push_listener(reason=f"unsupported:{reason}", force=True) + return + + self._cancel_backend_push_listener(reason=f"restart:{reason}", force=True) + + async def _listen() -> None: + try: + while not self.finalized.is_set(): + if user_seq != self.state.user_final_seq: + return + + pipeline = self._agent.pipeline + if pipeline is None: + return + + wait_for_push = getattr(pipeline, "wait_for_server_push", None) + if not callable(wait_for_push): + return + + raw_reply = await wait_for_push() + if raw_reply is None or self.finalized.is_set(): + return + if user_seq != self.state.user_final_seq: + return + + self._record_inflight_backend_activity() + reply = self._coerce_backend_reply( + raw_reply, + fallback_stage=self.state.current_stage, + ) + self._update_backend_processing_drop_from_reply(reply, source="backend_push") + self.state.current_stage = reply.stage + if not reply.text: + self._log_backend_reply_error_event(reply, source="backend_push") + + if reply.text: + if not await self._wait_for_stable_user_silence_before_terminal_reply( + reply, + user_seq=user_seq, + source="backend_push", + ): + return + try: + await self._speak_backend_reply( + reply, + source="backend_push", + add_to_chat_ctx=True, + ) + except Exception: + self._logger.exception("[tts] failed speaking backend push reply") + if self._timeline is not None: + self._timeline.emit( + "tts_stage_failed", + stage=self.state.current_stage, + source="backend_push", + ) + await self.terminate_with_resource_stop( + resource="tts", + source="backend_push_tts", + ) + return + + self.arm_agent_wait_timeout_from_reply( + reply, + user_seq=user_seq, + source="backend_push", + ) + + terminal_stop = self._terminal_stop_from_reply(reply) + if terminal_stop is not None: + await self._finish_with_agent_final_stop( + terminal_stop, + source="backend_push", + ) + return + + if reply.done or self._finalization_policy.is_done_stage(reply.stage): + self._create_task_logged( + self.notify_bridge_done("stage_done"), + name="bridge_done_push", + ) + self._create_task_logged( + self.finalize("stage_done"), + name="finalize_push", + ) + return + except asyncio.CancelledError: + return + except Exception: + self._logger.exception("[pipeline] backend push listener failed") + if self._timeline is not None: + self._timeline.emit("backend_push_failed", reason=reason) + await self.terminate_with_resource_stop( + resource="agent_backend", + source="backend_push_listener", + ) + + self._backend_push_task = self._create_task_logged( + _listen(), + name=f"backend_push_{user_seq}", + ) + if self._timeline is not None: + self._timeline.emit( + "backend_push_started", + reason=reason, + user_seq=user_seq, + inflight=inflight, + ) + + def _coerce_backend_reply(self, raw_reply: Any, *, fallback_stage: str) -> BackendReply: + if isinstance(raw_reply, BackendReply): + stage = (raw_reply.stage or fallback_stage or "UNKNOWN").strip().upper() + done = bool(raw_reply.done) or stage == "DONE" + return BackendReply( + stage=stage, + text=(raw_reply.text or "").strip(), + done=done, + export_payload=raw_reply.export_payload, + metadata=raw_reply.metadata, + ) + + if isinstance(raw_reply, tuple) and len(raw_reply) == 2: + stage_raw, text_raw = raw_reply + stage = str(stage_raw or fallback_stage or "UNKNOWN").strip().upper() + return BackendReply( + stage=stage, + text=str(text_raw or "").strip(), + done=stage == "DONE", + export_payload=None, + ) + + if isinstance(raw_reply, dict): + stage = str(raw_reply.get("stage") or fallback_stage or "UNKNOWN").strip().upper() + text = str(raw_reply.get("text") or "").strip() + done = bool(raw_reply.get("done")) or stage == "DONE" + export_payload = raw_reply.get("export_payload") + if export_payload is None and "result" in raw_reply: + export_payload = raw_reply.get("result") + return BackendReply( + stage=stage, + text=text, + done=done, + export_payload=export_payload, + metadata=raw_reply.get("metadata"), + ) + + stage = str(fallback_stage or self.state.current_stage or "UNKNOWN").strip().upper() + return BackendReply( + stage=stage, + text="", + done=stage == "DONE", + export_payload=raw_reply, + ) + + @staticmethod + def _wait_timeout_from_reply(reply: BackendReply) -> float: + metadata = reply.metadata + if not isinstance(metadata, Mapping): + return 0.0 + + raw_timeout = metadata.get("wait_timeout_seconds") + try: + timeout_s = float(raw_timeout) + except (TypeError, ValueError): + return 0.0 + + if timeout_s <= 0: + return 0.0 + return timeout_s + + @staticmethod + def _reply_has_wait_timeout_metadata(reply: BackendReply) -> bool: + metadata = reply.metadata + return isinstance(metadata, Mapping) and "wait_timeout_seconds" in metadata + + @staticmethod + def _wait_retry_messages_from_reply(reply: BackendReply) -> tuple[str, ...]: + metadata = reply.metadata + if not isinstance(metadata, Mapping): + return () + + raw_messages = metadata.get("wait_retry_messages") + if not isinstance(raw_messages, (list, tuple)): + return () + + messages: list[str] = [] + for raw_message in raw_messages: + message = str(raw_message or "").strip() + if message: + messages.append(message) + return tuple(messages) + + @classmethod + def _wait_timeout_config_from_reply( + cls, + reply: BackendReply, + ) -> AgentWaitTimeoutConfig | None: + timeout_s = cls._wait_timeout_from_reply(reply) + if timeout_s <= 0: + return None + return AgentWaitTimeoutConfig( + timeout_s=timeout_s, + retry_messages=cls._wait_retry_messages_from_reply(reply), + ) + + @staticmethod + def _terminal_stop_from_reply(reply: BackendReply) -> AgentFinalStop | None: + return final_stop_from_agent_result(reply.export_payload) + + async def _wait_for_stable_user_silence_before_terminal_reply( + self, + reply: BackendReply, + *, + user_seq: int, + source: str, + ) -> bool: + """Do not let a terminal reply start on top of an active user utterance. + + User state changes to ``listening`` slightly before the final transcript is + delivered. Requiring a short stable-silence window keeps the terminal + message from starting on top of the user, but a late transcript must not + invalidate a reply that closes the conversation. + """ + terminal = self._terminal_stop_from_reply(reply) is not None or ( + reply.done or self._finalization_policy.is_done_stage(reply.stage) + ) + if not terminal: + return True + + waited = False + while not self.finalized.is_set(): + if not self._user_not_speaking.is_set(): + if not waited: + waited = True + log_flow_event( + self._call_logger, + "terminal_tts_wait_user_silence", + source=source, + stage=reply.stage, + user_seq=user_seq, + ) + await self._user_not_speaking.wait() + + await asyncio.sleep(POST_USER_FINAL_SPEECH_GRACE_S) + if self._user_not_speaking.is_set(): + break + + allowed = not self.finalized.is_set() + if not allowed: + log_flow_event( + self._call_logger, + "terminal_tts_skipped", + reason="already_finalized", + source=source, + stage=reply.stage, + user_seq=user_seq, + current_seq=self.state.user_final_seq, + ) + elif waited: + log_flow_event( + self._call_logger, + "terminal_tts_user_silence_confirmed", + source=source, + stage=reply.stage, + user_seq=user_seq, + ) + return allowed + + async def _finish_with_agent_final_stop( + self, + stop: AgentFinalStop, + *, + source: str, + ) -> None: + async with self._terminal_stop_lock: + if self._terminal_stop_sent: + return + self._terminal_stop_sent = True + + if self._timeline is not None: + self._timeline.emit( + "agent_final_stop_started", + status=stop.status, + reason=stop.reason, + result_type=stop.result_type, + source=source, + ) + + if source == "backend_push": + self._backend_push_task = None + else: + self._cancel_backend_push_listener(reason=f"agent_final_stop:{stop.reason}", force=True) + self.cancel_idle_timer("agent_final_stop") + self.cancel_idle_close("agent_final_stop") + self.cancel_agent_wait_timeout("agent_final_stop") + + self._log_resource_error_events( + status=stop.status, + reason=stop.reason, + source=source, + terminal=True, + ) + await self.notify_bridge_stop( + status=stop.status, + reason=stop.reason, + phase=stop.phase, + ) + self.finalized.set() + if self._timeline is not None: + self._timeline.emit( + "agent_final_stop_completed", + status=stop.status, + reason=stop.reason, + result_type=stop.result_type, + source=source, + ) + + async def _speak_backend_reply( + self, + reply: BackendReply, + *, + initial_turn: bool = False, + source: str, + add_to_chat_ctx: bool, + ) -> None: + structured_erro_msg = self._error_message_from_backend_reply(reply) + message_id = self._message_id_from_reply(reply) + agent_message_type = self._reply_agent_message_type(reply) + agent_result_type = self._reply_agent_result_type(reply) + if ( + agent_message_type == "feedback" + or agent_result_type == "feedback" + or self._metadata_str(reply.metadata, "event") == "feedback" + ): + message_id = self._feedback_message_id( + message_id=message_id, + speech_id=self._speech_id_from_reply(reply), + ) + self._update_backend_processing_drop_from_reply(reply, source=source) + if not reply.text: + self._log_backend_reply_error_event(reply, source=source, message_id=message_id) + return + self._call_logger.debug( + "BACKEND_REPLY_TEXT | source=%s | stage=%s | done=%s | text=%r", + source, + reply.stage, + reply.done, + reply.text, + ) + if self.finalized.is_set(): + self._call_logger.debug( + "BACKEND_TTS_SKIPPED | source=%s | stage=%s | reason=already_finalized", + source, + reply.stage, + ) + log_flow_event( + self._call_logger, + "tts_skip", + source=source, + stage=reply.stage, + reason="already_finalized", + ) + return + if not self._has_remote_participants(): + self._call_logger.debug( + "BACKEND_TTS_SKIPPED | source=%s | stage=%s | reason=no_remote_participants", + source, + reply.stage, + ) + log_flow_event( + self._call_logger, + "tts_skip", + source=source, + stage=reply.stage, + reason="no_remote_participants", + ) + return + + default_allow_interruptions = self._interrupt_policy.allow_stage(reply.stage) + allow_interruptions = self._allow_interruptions_for_reply( + reply, + default_allow_interruptions, + ) + if self._terminal_stop_from_reply(reply) is not None or ( + reply.done or self._finalization_policy.is_done_stage(reply.stage) + ): + # A resposta terminal e a confirmacao audivel do desfecho da chamada. + # Ela pode esperar o usuario terminar, mas nunca pode ser cortada. + allow_interruptions = False + if self.state.deferred_interruption.replay_in_flight: + # A resposta do reenvio nao pode ser interrompida. Descartar a fala do + # cliente no runtime nao basta: o barge-in do proprio LiveKit corta o + # playout por atividade de audio, e R2 sairia pela metade. + allow_interruptions = False + expects_user_response = self._reply_expects_user_response(reply) + initial_audio = None + if initial_turn: + cache_lookup = getattr(self._initial_greeting_tts, "initial_greeting_audio", None) + if callable(cache_lookup): + try: + initial_audio = cache_lookup(reply.text) + except Exception: + self._logger.exception("INITIAL_GREETING_AUDIO_CACHE_LOOKUP_FAIL") + if self._timeline is not None: + self._timeline.emit( + "initial_greeting_audio_cache", + outcome="hit" if initial_audio is not None else "miss_or_disabled", + ) + await self.say_stage( + reply.text, + reply.stage, + add_to_chat_ctx=add_to_chat_ctx, + allow_interruptions=allow_interruptions, + speech_id=self._speech_id_from_reply(reply), + message_id=message_id, + schedule_idle_after=True if expects_user_response is None else expects_user_response, + drop_user_input_while_speaking=self._reply_drops_user_input_while_speaking(reply), + audio=initial_audio, + agent_message_type=agent_message_type, + agent_result_type=agent_result_type, + erro_msg=structured_erro_msg, + ) + + def cancel_agent_wait_timeout(self, reason: str = "cancel") -> None: + self._scheduler.cancel( + "agent_wait_timeout", + reason=reason, + log_name="AGENT_WAIT_TIMEOUT_CANCEL", + ) + + def _remember_agent_wait_timeout( + self, + config: AgentWaitTimeoutConfig, + *, + user_seq: int, + source: str, + message_id: str, + ) -> None: + self._last_agent_wait_timeout_config = config + self._last_agent_wait_timeout_user_seq = user_seq + self._last_agent_wait_timeout_source = source + self._last_agent_wait_timeout_message_id = message_id + + def _clear_agent_wait_timeout_memory(self) -> None: + self._last_agent_wait_timeout_config = None + self._last_agent_wait_timeout_user_seq = 0 + self._last_agent_wait_timeout_source = "" + self._last_agent_wait_timeout_message_id = "" + + def _clear_agent_wait_timeout_config(self) -> None: + self._agent_wait_timeout_config = None + + def _restore_agent_wait_timeout_after_empty_stt(self) -> bool: + config = self._last_agent_wait_timeout_config + if config is None: + return False + if self._last_agent_wait_timeout_user_seq != self.state.user_final_seq: + return False + if self.finalized.is_set() or (self.state.current_stage or "").upper() == "DONE": + return False + + previous_source = self._last_agent_wait_timeout_source + previous_message_id = self._last_agent_wait_timeout_message_id + self._arm_agent_wait_timeout( + config, + user_seq=self.state.user_final_seq, + source="empty_stt_final", + message_id=previous_message_id, + ) + log_flow_event( + self._call_logger, + "stt_final_empty_recovery", + action="rearm_agent_wait_timeout", + previous_source=previous_source, + message_id=previous_message_id, + user_seq=self.state.user_final_seq, + stage=self.state.current_stage, + ) + return True + + def _recover_after_empty_stt_final(self) -> None: + if self.finalized.is_set(): + self._empty_stt_recovery_pending = False + return + if self.speaking.is_set() or self.state.gap.active: + if not self._pre_backend_wait_notice_is_current_speech(): + self._cancel_pre_backend_wait_notice(reason="empty_stt_final") + self._empty_stt_recovery_pending = True + log_flow_event( + self._call_logger, + "stt_final_empty_recovery", + action="defer_while_agent_active", + stage=self.state.current_stage, + speaking=self.speaking.is_set(), + gap_active=self.state.gap.active, + ) + return + self._cancel_pre_backend_wait_notice(reason="empty_stt_final") + self._empty_stt_recovery_pending = False + if self._restore_agent_wait_timeout_after_empty_stt(): + return + + self.arm_idle_timer(reason="empty_stt_final") + log_flow_event( + self._call_logger, + "stt_final_empty_recovery", + action="arm_idle_timer", + stage=self.state.current_stage, + ) + + def handle_empty_stt_final(self, *, transcript_len: int = 0, source: str = "stt_final") -> None: + deferred = self.state.deferred_interruption + if source != "user_input_transcribed": + # O provider avisa o final vazio antes de devolver o SpeechEvent. Em + # raw_json o transcript cru nao e vazio, o LiveKit repassa o evento e + # a mesma fala chegaria duas vezes ate a fila VAD. + self._deferred_empty_final_latched_at = time.monotonic() + if deferred.replay_in_flight: + log_flow_event(self._call_logger, "deferred_replay_stt_discarded", text="") + return + if self._handle_deferred_stt_final(transcription="", text=""): + return + + log_flow_event( + self._call_logger, + "stt_final_empty", + reason="empty_text", + source=source, + stage=self.state.current_stage, + speaking=self.speaking.is_set(), + transcript_len=transcript_len, + pre_backend_wait_reserved=self._pre_backend_wait_notice_reserved, + ) + if self._timeline is not None: + self._timeline.emit( + "user_transcript_final_empty", + current_stage=self.state.current_stage, + speaking=self.speaking.is_set(), + transcript_len=transcript_len, + pre_backend_wait_reserved=self._pre_backend_wait_notice_reserved, + ) + self._recover_after_empty_stt_final() + + async def _stop_due_to_agent_wait_timeout( + self, + *, + reason: str, + source: str, + message_id: str = "", + ) -> None: + async with self._terminal_stop_lock: + if self._terminal_stop_sent: + return + self._terminal_stop_sent = True + + self.cancel_idle_timer("agent_wait_timeout") + self.cancel_idle_close("agent_wait_timeout") + if self._timeline is not None: + self._timeline.emit( + "agent_wait_timeout_fired", + status=AGENT_WAIT_TIMEOUT_STOP_STATUS, + reason=reason, + source=source, + message_id=message_id, + ) + + self._restore_agent_wait_timeout_retry_vad_threshold( + reason=f"agent_wait_timeout:{reason}" + ) + await self.notify_bridge_stop( + status=AGENT_WAIT_TIMEOUT_STOP_STATUS, + reason=reason, + phase="in_session", + ) + self._log_resource_error_events( + status=AGENT_WAIT_TIMEOUT_STOP_STATUS, + reason=reason, + source=source, + message_id=message_id, + terminal=True, + ) + self.finalized.set() + + def _arm_agent_wait_timeout( + self, + config: AgentWaitTimeoutConfig, + *, + user_seq: int, + source: str, + message_id: str = "", + ) -> None: + timeout_s = config.timeout_s + if timeout_s <= 0: + self.cancel_agent_wait_timeout(f"clear:{source}") + self._clear_agent_wait_timeout_memory() + return + + timeout_message_id = self._resolve_agent_message_id(message_id, generate=True) + self._remember_agent_wait_timeout( + config, + user_seq=user_seq, + source=source, + message_id=timeout_message_id, + ) + retry_messages = config.retry_messages + self.cancel_idle_timer("agent_wait_timeout") + self.cancel_idle_close("agent_wait_timeout") + self.cancel_agent_wait_timeout(f"rearm:{source}") + reason = self._idle_policy.end_reason + token = 0 + + def _should_continue() -> bool: + if not self._scheduler.is_current("agent_wait_timeout", token): + return False + if user_seq != self.state.user_final_seq: + return False + if self.finalized.is_set(): + return False + return True + + async def _fire() -> None: + for attempt, retry_text in enumerate(retry_messages, start=1): + try: + await asyncio.sleep(timeout_s) + except asyncio.CancelledError: + return + + if not _should_continue(): + return + + self.cancel_idle_timer("agent_wait_timeout_retry") + self.cancel_idle_close("agent_wait_timeout_retry") + inactivity_message_id = self._agent_wait_timeout_inactivity_message_id( + message_id=timeout_message_id, + attempt=attempt, + ) + self._call_logger.info( + "AGENT_WAIT_TIMEOUT_RETRY | attempt=%s/%s | wait_s=%.1f | message_id=%s | text=%r", + attempt, + len(retry_messages), + timeout_s, + inactivity_message_id, + retry_text, + ) + if self._timeline is not None: + self._timeline.emit( + "agent_wait_timeout_retry", + wait_timeout_s=timeout_s, + attempt=attempt, + max_attempts=len(retry_messages), + reason=reason, + source=source, + user_seq=user_seq, + text=retry_text, + message_id=inactivity_message_id, + ) + + self._activate_agent_wait_timeout_retry_vad_threshold( + attempt=attempt, + reason="agent_wait_timeout_retry", + ) + try: + if self._agent.pipeline is not None: + await self.execute(InjectIdleNudge(retry_text)) + except Exception: + self._logger.debug("IDLE_NUDGE_PIPELINE_HOOK_FAIL", exc_info=True) + + try: + await self.say_stage( + retry_text, + "AGENT_WAIT_TIMEOUT_RETRY", + add_to_chat_ctx=False, + schedule_idle_after=False, + message_id=inactivity_message_id, + ) + except Exception: + self._logger.exception("[tts] failed speaking agent wait retry") + if self._timeline is not None: + self._timeline.emit( + "tts_stage_failed", + stage="AGENT_WAIT_TIMEOUT_RETRY", + source="agent_wait_timeout", + message_id=inactivity_message_id, + ) + await self.terminate_with_resource_stop( + resource="tts", + source="agent_wait_timeout_retry_tts", + ) + return + + if not _should_continue(): + return + + try: + await asyncio.sleep(timeout_s) + except asyncio.CancelledError: + return + + if not _should_continue(): + return + await self._stop_due_to_agent_wait_timeout( + reason=reason, + source=source, + message_id=timeout_message_id, + ) + + token = self._scheduler.arm( + "agent_wait_timeout", + task_name="agent_wait_timeout_timer", + coro=_fire(), + ) + if self._timeline is not None: + self._timeline.emit( + "agent_wait_timeout_armed", + wait_timeout_s=timeout_s, + retry_messages=len(retry_messages), + reason=reason, + source=source, + user_seq=user_seq, + ) + + def arm_agent_wait_timeout_from_reply( + self, + reply: BackendReply, + *, + user_seq: int, + source: str, + ) -> None: + stage = str(reply.stage or self.state.current_stage or "").strip().upper() + if reply.done or self._finalization_policy.is_done_stage(stage): + self.cancel_agent_wait_timeout(f"done:{source}") + self._clear_agent_wait_timeout_memory() + self._clear_agent_wait_timeout_config() + return + + expects_user_response = self._reply_expects_user_response(reply) + if expects_user_response is False: + self.cancel_agent_wait_timeout(f"not_expected:{source}") + self._clear_agent_wait_timeout_memory() + return + + incoming_config = self._wait_timeout_config_from_reply(reply) + has_wait_timeout_metadata = self._reply_has_wait_timeout_metadata(reply) + if incoming_config is not None: + self._agent_wait_timeout_config = incoming_config + elif has_wait_timeout_metadata: + self.cancel_agent_wait_timeout(f"disable:{source}") + self._clear_agent_wait_timeout_memory() + self._clear_agent_wait_timeout_config() + return + elif not reply.text: + self._clear_agent_wait_timeout_memory() + return + + config = self._agent_wait_timeout_config + if config is None: + self._clear_agent_wait_timeout_memory() + return + + self._arm_agent_wait_timeout( + config, + user_seq=user_seq, + source=source, + message_id=self._message_id_from_reply(reply), + ) + + async def notify_bridge_done(self, reason: str = "stage_done") -> None: + try: + await self.execute(NotifyBridgeDone(reason)) + if self._timeline is not None: + self._timeline.emit("bridge_done_requested", reason=reason) + except Exception as exc: + self._log_resource_error_events( + status="stop_bridge_failed", + reason="bridge_notify_done_failed", + resource="bridge", + source="notify_bridge_done", + exc=exc, + ) + self._logger.exception("[bridge] Falha enviando DONE") + + async def finalize(self, reason: str) -> None: + if self._finalization_policy.should_skip(finalized=self.finalized.is_set()): + return + + async with self.finalize_lock: + if self._finalization_policy.should_skip(finalized=self.finalized.is_set()): + return + self._restore_agent_wait_timeout_retry_vad_threshold( + reason=f"finalize:{reason}" + ) + if self._timeline is not None: + self._timeline.emit("finalize_started", reason=reason) + + self._cancel_backend_push_listener(reason=f"finalize:{reason}", force=True) + self.cancel_idle_timer("finalize") + self.cancel_idle_close("finalize") + self.cancel_agent_wait_timeout("finalize") + self._reset_deferred_interruption() + self.release_user_audio_input(reason=f"finalize:{reason}") + + dur = time.monotonic() - self._call_t0 + self._call_logger.info( + "CALL_END | room=%s | protocol=%s | session_id=%s | reason=%s | duration_s=%.1f", + self._ctx.room.name, + self._protocol, + self._session_id, + reason, + dur, + ) + + try: + await asyncio.wait_for(self._agent._ready.wait(), timeout=5.0) + except Exception: + pass + + try: + raw_reply = await self.execute(EndServiceOnce()) + reply = self._coerce_backend_reply( + raw_reply, + fallback_stage=self.state.current_stage, + ) + self.state.current_stage = reply.stage + await self._speak_backend_reply( + reply, + source="end_service_once", + add_to_chat_ctx=False, + ) + export_payload = reply.export_payload if reply.export_payload is not None else [] + await self.execute(ExportSession(output=export_payload, session_id=self._session_id)) + except Exception: + self._logger.exception( + "[finalize] Falha ao finalizar/exportar session_id=%s", + self._session_id, + ) + + self.finalized.set() + if self._timeline is not None: + self._timeline.emit("finalize_completed", reason=reason) + + def cancel_idle_timer(self, reason: str = "cancel") -> None: + self._scheduler.cancel("idle", reason=reason, log_name="IDLE_TIMER_CANCEL") + + def cancel_idle_close(self, reason: str = "cancel") -> None: + self._scheduler.cancel("idle_close", reason=reason, log_name="IDLE_CLOSE_CANCEL") + + def note_user_speech_activity(self, *, reason: str) -> None: + self.cancel_idle_timer(reason) + self.cancel_idle_close(reason) + self.cancel_agent_wait_timeout(reason) + self.state.idle_nudge_count = 0 + if self._timeline is not None: + self._timeline.emit( + "user_speech_activity_detected", + reason=reason, + speaking=self.speaking.is_set(), + current_stage=self.state.current_stage, + ) + + async def _end_due_to_no_response(self) -> None: + end_reason = self._idle_policy.end_reason + self._restore_agent_wait_timeout_retry_vad_threshold(reason="idle_no_response") + self._log_resource_error_events( + status=AGENT_WAIT_TIMEOUT_STOP_STATUS, + reason=end_reason, + source="idle_no_response", + terminal=True, + ) + self._create_task_logged( + self.notify_bridge_done(end_reason), + name="bridge_done_noresp", + ) + self._create_task_logged( + self.finalize(end_reason), + name="finalize_noresp", + ) + + def arm_idle_close_timer(self, *, reason: str) -> None: + if self._idle_policy.close_delay_s <= 0: + self._create_task_logged( + self._end_due_to_no_response(), + name="idle_close_immediate", + ) + return + + self.cancel_idle_timer("close_arm") + self.cancel_idle_close(f"rearm:{reason}") + seq_snapshot = self.state.user_final_seq + token = 0 + + async def _close() -> None: + try: + await asyncio.sleep(self._idle_policy.close_delay_s) + except asyncio.CancelledError: + return + + if not self._idle_policy.should_fire_close( + token_is_current=self._scheduler.is_current("idle_close", token), + seq_matches=self.state.user_final_seq == seq_snapshot, + finalized=self.finalized.is_set(), + current_stage=self.state.current_stage, + ): + return + + t0 = time.monotonic() + while ( + (self.speaking.is_set() or self.state.gap.active) + and not self.finalized.is_set() + and (time.monotonic() - t0) < 3.0 + ): + await asyncio.sleep(0.25) + if not self._scheduler.is_current("idle_close", token): + return + if self.state.user_final_seq != seq_snapshot: + return + + self._call_logger.info( + "IDLE_CLOSE_FIRE | after_max_tries=%s | wait_s=%.1f | reason=%s", + self._idle_policy.max_tries, + self._idle_policy.close_delay_s, + self._idle_policy.end_reason, + ) + if self._timeline is not None: + self._timeline.emit( + "idle_close_fire", + reason=self._idle_policy.end_reason, + max_tries=self._idle_policy.max_tries, + ) + await self._end_due_to_no_response() + + token = self._scheduler.arm( + "idle_close", + task_name="idle_close_timer", + coro=_close(), + ) + + def arm_idle_timer(self, *, reason: str, delay_s: Optional[float] = None) -> None: + delay = self._idle_policy.resolve_delay(delay_s) + if not self._idle_policy.should_arm_nudge(delay_s=delay, nudge_text=self._nudge_text): + return + + self.cancel_idle_timer(f"rearm:{reason}") + seq_snapshot = self.state.user_final_seq + token = 0 + + async def _timer() -> None: + try: + await asyncio.sleep(delay) + except asyncio.CancelledError: + return + + if not self._idle_policy.should_fire_nudge( + token_is_current=self._scheduler.is_current("idle", token), + seq_matches=self.state.user_final_seq == seq_snapshot, + speaking=self.speaking.is_set(), + gap_active=self.state.gap.active, + finalized=self.finalized.is_set(), + current_stage=self.state.current_stage, + ): + return + + if self._idle_policy.should_arm_close_before_fire(nudge_count=self.state.idle_nudge_count): + self._call_logger.info( + "IDLE_NUDGE_MAX_REACHED | tries=%s | will_close_in_s=%.1f", + self.state.idle_nudge_count, + self._idle_policy.close_delay_s, + ) + self.arm_idle_close_timer(reason="max_reached_before_fire") + return + + self.state.idle_nudge_count += 1 + self._call_logger.info( + "IDLE_NUDGE_FIRE | attempt=%s/%s | delay_s=%.1f | text=%r", + self.state.idle_nudge_count, + self._idle_policy.max_tries, + delay, + self._nudge_text, + ) + if self._timeline is not None: + self._timeline.emit( + "idle_nudge_fire", + attempt=self.state.idle_nudge_count, + max_tries=self._idle_policy.max_tries, + text=self._nudge_text, + ) + + try: + if self._agent.pipeline is not None: + await self.execute(InjectIdleNudge(self._nudge_text)) + except Exception: + self._logger.debug("IDLE_NUDGE_PIPELINE_HOOK_FAIL", exc_info=True) + + inactivity_message_id = self._idle_nudge_message_id( + message_id=self._current_base_message_id(), + attempt=self.state.idle_nudge_count, + ) + await self.say_stage( + self._nudge_text, + "IDLE_NUDGE", + add_to_chat_ctx=False, + message_id=inactivity_message_id, + ) + + token = self._scheduler.arm( + "idle", + task_name="idle_timer", + coro=_timer(), + ) + + def should_ignore_backchannel(self, text: str) -> bool: + return False + + def _remember_interrupt_task(self, task: asyncio.Task[Any]) -> asyncio.Task[Any]: + self.state.current_speech.interruption_task = task + self._last_interrupt_task = task + return task + + def request_current_speech_interruption( + self, + *, + reason: str, + force: bool = False, + ) -> asyncio.Task[Any] | None: + if not force and not self._current_speech_allows_interruption(): + return None + + self._mark_structured_interruption(reason=reason) + existing = self.state.current_speech.interruption_task + if self.state.current_speech.interruption_requested and existing is not None: + return existing + + self.state.current_speech.interruption_requested = True + + handle = self.state.current_speech.handle + speech_id = self.state.current_speech.speech_id + stage = self.state.current_speech.stage + + self._call_logger.info( + "BARGE_IN_INTERRUPT_REQUEST | reason=%s | stage=%s | speech_id=%s | has_handle=%s", + reason, + stage or "-", + speech_id or "-", + handle is not None, + ) + if self._timeline is not None: + self._timeline.emit( + "barge_in_interrupt_requested", + reason=reason, + stage=stage or "", + speech_id=speech_id, + has_handle=handle is not None, + ) + + async def _interrupt() -> None: + if handle is None: + return + await self.execute(InterruptSpeech(handle, force=force)) + + return self._remember_interrupt_task( + self._create_task_logged( + _interrupt(), + name=f"interrupt_current_speech_{reason}", + ) + ) + + @staticmethod + def _track_publication_fields(publication: Any) -> dict[str, Any]: + source = getattr(publication, "source", "") + source_name = getattr(source, "name", None) or str(source) + track = getattr(publication, "track", None) + return { + "track_sid": str( + getattr(publication, "sid", "") + or getattr(track, "sid", "") + or "" + ), + "track_name": str(getattr(publication, "name", "") or getattr(track, "name", "") or ""), + "source": source_name, + "kind": str(getattr(publication, "kind", "") or getattr(track, "kind", "") or ""), + "muted": bool(getattr(publication, "muted", False)), + "subscribed": bool(getattr(publication, "subscribed", False)), + "has_track": track is not None, + } + + def log_room_input_snapshot(self, *, reason: str) -> None: + room = getattr(self._ctx, "room", None) + participants = getattr(room, "remote_participants", {}) or {} + bridge_identity = (self._bridge_identity or "").strip() + participant = participants.get(bridge_identity) if bridge_identity else None + + if participant is None and len(participants) == 1: + participant = next(iter(participants.values())) + + if participant is None: + log_flow_event( + self._call_logger, + "agent_input_snapshot", + reason=reason, + bridge_identity=bridge_identity, + participant_found=False, + remote_participants=len(participants), + ) + return + + publications = list((getattr(participant, "track_publications", {}) or {}).values()) + audio_publications = [ + publication + for publication in publications + if "audio" in str(getattr(publication, "kind", "") or "").lower() + or "microphone" in str(getattr(publication, "source", "") or "").lower() + ] + selected = audio_publications[0] if audio_publications else (publications[0] if publications else None) + fields = self._track_publication_fields(selected) if selected is not None else {} + log_flow_event( + self._call_logger, + "agent_input_snapshot", + reason=reason, + bridge_identity=bridge_identity, + participant_found=True, + participant=str(getattr(participant, "identity", "") or ""), + remote_participants=len(participants), + publication_count=len(publications), + audio_publication_count=len(audio_publications), + **fields, + ) + + def mark_interrupt_from_handle_if_any(self, *, reason: str) -> asyncio.Task[Any]: + handle = self.state.current_speech.handle + speech_id = self.state.current_speech.speech_id + if handle is None: + task = self._create_task_logged( + self.execute( + SetPendingInterrupt( + listened_text="", + skipped=True, + speech_id=speech_id, + ) + ), + name="pending_interrupt_nohandle", + ) + self._call_logger.info( + "INTERRUPT_MARK | reason=%s | skipped=True | speech_id=%s | spoken=''", + reason, + speech_id or "-", + ) + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="no_current_speech_handle", + source=reason, + stage=self.state.current_speech.stage or self.state.current_stage, + skipped=True, + speech_id=speech_id, + ) + if self._timeline is not None: + self._timeline.emit( + "interrupt_marked", + reason=reason, + skipped=True, + spoken="", + speech_id=speech_id, + ) + self.state.current_speech.interruption_requested = True + return self._remember_interrupt_task(task) + + spoken = self.execute_now(ExtractSpokenText(handle)) + skipped = not bool((spoken or "").strip()) + task = self._create_task_logged( + self.execute( + SetPendingInterrupt( + listened_text=spoken, + skipped=skipped, + speech_id=speech_id, + ) + ), + name="pending_interrupt_handle", + ) + self._call_logger.info( + "INTERRUPT_MARK | reason=%s | stage=%s | skipped=%s | speech_id=%s | spoken=%r", + reason, + self.state.current_speech.stage or "-", + skipped, + speech_id or "-", + spoken, + ) + log_flow_event( + self._call_logger, + "interrupt_marked", + reason=reason, + stage=self.state.current_speech.stage or self.state.current_stage, + skipped=skipped, + text=spoken, + speech_id=speech_id, + ) + if skipped: + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="empty_spoken_text", + source=reason, + stage=self.state.current_speech.stage or self.state.current_stage, + skipped=True, + speech_id=speech_id, + ) + if self._timeline is not None: + self._timeline.emit( + "interrupt_marked", + reason=reason, + stage=self.state.current_speech.stage or "", + skipped=skipped, + spoken=spoken, + speech_id=speech_id, + ) + self.state.current_speech.interruption_requested = True + return self._remember_interrupt_task(task) + + async def say_stage( + self, + text: str, + stage: str, + *, + add_to_chat_ctx: bool = True, + allow_interruptions: Optional[bool] = None, + rearm_interruptions_on_next_user_speech: bool = False, + speech_id: str = "", + message_id: str = "", + schedule_idle_after: bool = True, + wait_post_final_grace: bool = True, + require_user_not_speaking: bool = False, + abort_if_agent_spoke_since: int | None = None, + audio: Any = None, + audio_duration_ms: int | None = None, + drop_user_input_while_speaking: bool = False, + agent_message_type: str = "", + agent_result_type: str = "", + erro_msg: str | None = None, + playout_result_callback: Optional[Callable[..., None]] = None, + ) -> bool: + text = (text or "").strip() + if not text: + return False + message_id = str(message_id or "").strip() + stage_key = str(stage or "").strip().upper() + if stage_key not in AGENT_WAIT_TIMEOUT_RETRY_VAD_STAGES: + self._restore_agent_wait_timeout_retry_vad_threshold( + reason=f"agent_stage:{stage_key or 'UNKNOWN'}" + ) + if require_user_not_speaking and not self._user_not_speaking.is_set(): + log_flow_event( + self._call_logger, + "tts_skipped", + stage=stage, + reason="user_speaking", + message_id=message_id or self._current_base_message_id(), + text=text, + ) + return False + + self._clear_post_playout_drop_window(reason="new_tts") + allow_int = ( + self._interrupt_policy.allow_stage(stage) + if allow_interruptions is None + else bool(allow_interruptions) + ) + speech_id = str(speech_id or "").strip() + message_id = self._resolve_agent_message_id(message_id, generate=True) + + self.cancel_idle_timer("bot_speaking") + self.cancel_idle_close("bot_speaking") + + if self.state.gap.active and self.state.gap.cancelled: + self._call_logger.info("BOT_SAY_SKIPPED | stage=%s | reason=gap_cancelled", stage) + self.reset_gap_state() + return False + + if self.state.gap.active and self.state.user_final_seq != self.state.gap.guard_seq: + await self.execute( + SetPendingInterrupt( + listened_text="", + skipped=True, + speech_id=self.state.gap.speech_id, + ) + ) + self._call_logger.info("BOT_SAY_SKIPPED | stage=%s | reason=user_spoke_in_gap", stage) + self.reset_gap_state() + return False + + async with self.say_lock: + if require_user_not_speaking and not self._user_not_speaking.is_set(): + log_flow_event( + self._call_logger, + "tts_skipped", + stage=stage, + reason="user_speaking_after_lock", + message_id=message_id, + text=text, + ) + return False + if ( + abort_if_agent_spoke_since is not None + and self._say_stage_seq != abort_if_agent_spoke_since + ): + # Outra fala ocupou o say_lock enquanto este audio esperava na + # fila: o silencio que ele mascararia ja foi preenchido. + log_flow_event( + self._call_logger, + "tts_skipped", + stage=stage, + reason="agent_spoke_after_schedule", + message_id=message_id, + text=text, + ) + return False + self._say_stage_seq += 1 + self.state.current_speech.stage = stage + self.state.current_speech.text = text + self.state.current_speech.speech_id = speech_id + self.state.current_speech.drop_user_input_while_speaking = bool( + drop_user_input_while_speaking + ) + self.state.current_speech.agent_message_type = str(agent_message_type or "").strip() + self.state.current_speech.agent_result_type = str(agent_result_type or "").strip() + self.state.current_speech.allow_interruptions = allow_int + self.state.current_speech.rearm_interruptions_on_next_user_speech = bool( + rearm_interruptions_on_next_user_speech + ) + self.state.current_speech.interruption_requested = False + self.state.current_speech.interruption_task = None + self.speaking.set() + log_flow_event( + self._call_logger, + "tts_start", + stage=stage, + allow_interruptions=allow_int, + message_id=message_id, + text=text, + ) + if self._timeline is not None: + self._timeline.emit( + "tts_stage_started", + stage=stage, + allow_interruptions=allow_int, + text=text, + ) + self._emit_debug_event( + "agent.speech.started", + stage=stage, + message_id=message_id, + text=text, + ) + seq_at_start = self.state.user_final_seq + handle = None + tts_started_ns = None + tts_started_monotonic = 0.0 + agent_speaking_seq_at_start = 0 + tts_tffb_ms = None + tts_audio_duration_ms = None + tts_playout_completed = False + + def _tts_was_interrupted() -> bool: + return bool(getattr(handle, "interrupted", False)) or ( + self.state.user_final_seq != seq_at_start + ) + + try: + self.reset_gap_state() + if wait_post_final_grace: + await self._wait_post_user_final_grace(stage=stage) + + async def _start_speech_attempt(): + started_ns = time.time_ns() + started_monotonic = time.monotonic() + speaking_seq_at_start = self._agent_speaking_seq + started_handle = await self.execute( + StartSpeech( + text=text, + allow_interruptions=allow_int, + add_to_chat_ctx=add_to_chat_ctx, + audio=audio, + ) + ) + self.state.current_speech.handle = started_handle + + self._call_logger.info( + "BOT_SAY_HANDLE | stage=%s | speech_handle_id=%s | seq_at_start=%s | current_seq=%s", + stage, + self._speech_handle_id(started_handle) or "-", + seq_at_start, + self.state.user_final_seq, + ) + return ( + started_handle, + started_ns, + started_monotonic, + speaking_seq_at_start, + ) + + ( + handle, + tts_started_ns, + tts_started_monotonic, + agent_speaking_seq_at_start, + ) = await _start_speech_attempt() + + try: + should_abort_tts = False + abort_reason = "empty_frame_timeout" + retry_error: Exception | None = None + retryable_error_detail: str | None = None + partial_failure: _TTSPartialFailure | None = None + if audio is None: + playout_start_timeout_s = self._tts_playout_start_timeout_s() + try: + should_abort_tts = await self._wait_for_tts_playout_or_empty_frame_retry( + handle, + started_monotonic=tts_started_monotonic, + agent_speaking_seq_at_start=agent_speaking_seq_at_start, + timeout_s=playout_start_timeout_s, + stage=stage, + message_id=message_id, + ) + except Exception as exc: + if not self._is_retryable_tts_playout_error(exc): + raise + retry_error = exc + abort_reason = "playout_error" + should_abort_tts = True + if not should_abort_tts: + partial_failure = self._pop_tts_partial_failure(handle) + if partial_failure is not None: + abort_reason = "provider_partial_failure" + retryable_error_detail = partial_failure.detail + should_abort_tts = True + else: + retryable_error_detail = self._pop_retryable_tts_error_detail(handle) + observed_audio_duration_ms = self._resolve_tts_audio_duration_ms( + handle + ) + if ( + retryable_error_detail + and ( + observed_audio_duration_ms is None + or "xai tts partial audio failure" in retryable_error_detail + ) + and not bool(getattr(handle, "interrupted", False)) + and self.state.user_final_seq == seq_at_start + ): + abort_reason = "provider_error_no_audio_metric" + should_abort_tts = True + log_flow_event( + self._call_logger, + "tts_playout_missing_audio_after_error", + stage=stage, + message_id=message_id, + speech_handle_id=self._speech_handle_id(handle), + detail=retryable_error_detail, + ) + else: + await self.execute(WaitForSpeechPlayout(handle)) + + if should_abort_tts: + if not bool(getattr(handle, "interrupted", False)): + try: + await self.execute(InterruptSpeech(handle, force=True)) + except Exception as exc: + log_flow_event( + self._call_logger, + "tts_failure_interrupt_error", + stage=stage, + message_id=message_id, + speech_handle_id=self._speech_handle_id(handle), + error=type(exc).__name__, + ) + tts_total_ms = self._resolve_tts_provider_duration_ms(handle) + tts_audio_duration_ms = self._resolve_tts_audio_duration_ms(handle) + if partial_failure is not None: + tts_total_ms = ( + partial_failure.provider_synthesis_ms or tts_total_ms + ) + tts_audio_duration_ms = ( + partial_failure.pcm_duration_ms + or tts_audio_duration_ms + ) + if retry_error is not None: + retry_error_detail = f"{type(retry_error).__name__}: {retry_error}" + http_cod_status, http_cod_desc = _http_status_details_from_exception( + retry_error + ) + elif retryable_error_detail: + retry_error_detail = retryable_error_detail + http_cod_status, http_cod_desc = None, None + else: + retry_error_detail = ( + "TimeoutError: nenhum frame de audio do TTS foi observado " + f"em {round(playout_start_timeout_s * 1000)} ms" + ) + http_cod_status, http_cod_desc = None, None + tts_error_message_id = self._tts_error_message_id( + message_id=message_id + ) + tts_max_gap_ms = self._tts_failure_metric(retry_error_detail, "max_audio_delta_gap_ms") + tts_max_underrun_0ms = self._tts_failure_metric(retry_error_detail, "xai_underrun_estimado_ms") + tts_underflow_count = self._tts_failure_metric(retry_error_detail, "xai_micro_underflows") + tts_avg_underflow_ms = self._tts_failure_metric(retry_error_detail, "xai_avg_underrun_ms") + log_structured_event( + self._call_logger, + self._structured_log_context, + tipo_evento=EVENT_ENVIO_MSG, + message_id=tts_error_message_id, + inicio_ns=tts_started_ns, + fim_ns=time.time_ns(), + latencia_total_ms=tts_total_ms, + duracao_audio_ms=tts_audio_duration_ms, + tts_max_gap_ms=tts_max_gap_ms, + tts_max_underrun_0ms=tts_max_underrun_0ms, + tts_underflow_count=tts_underflow_count, + tts_avg_underflow_ms=tts_avg_underflow_ms, + interrupcao=True if _tts_was_interrupted() else None, + erro_msg="Falha TTS", + erro_detalhe=retry_error_detail, + http_cod_status=http_cod_status, + http_cod_desc=http_cod_desc, + ) + log_flow_event( + self._call_logger, + "tts_failure", + stage=stage, + message_id=tts_error_message_id, + reason=abort_reason, + erro_msg="Falha TTS", + erro_detalhe=retry_error_detail, + detail=retry_error_detail, + duration_ms=tts_total_ms, + audio_ms=tts_audio_duration_ms, + timeout_ms=round(playout_start_timeout_s * 1000), + timeout_kind="playout_start", + first_frame_timeout_ms=round(self._tts_first_frame_timeout_s() * 1000), + connection_grace_ms=round(self._tts_connection_start_grace_s() * 1000), + error=type(retry_error).__name__ if retry_error is not None else "", + ) + self._call_logger.warning( + "BOT_TTS_FAILURE | stage=%s | message_id=%s | reason=%s", + stage, + message_id or "-", + TTS_FAILURE_DETAIL, + ) + log_flow_event( + self._call_logger, + "tts_failure_return_to_listening", + stage=stage, + message_id=message_id, + reason=abort_reason, + socket_action="reconnect_on_next_synthesis", + action="wait_for_client_repetition", + ) + # The adapter has discarded the websocket after a partial failure. + # Do not replay a consumed LiveKit stream or add a recovery prompt. + return False + + tts_finished_ns = time.time_ns() + await self._resolve_tts_tffb_ms(handle) + ( + tts_total_ms, + tts_tffb_ms, + tts_audio_duration_ms, + tts_metric_duration_ms, + tts_metrics_source, + ) = self._resolve_tts_structured_metrics(handle) + tts_interrupted = _tts_was_interrupted() + ( + tts_max_gap_ms, + tts_max_underrun_0ms, + tts_underflow_count, + tts_avg_underflow_ms, + ) = self._resolve_tts_turn_health_metrics(handle) + log_structured_event( + self._call_logger, + self._structured_log_context, + tipo_evento=EVENT_ENVIO_MSG, + message_id=message_id, + inicio_ns=tts_started_ns, + fim_ns=tts_finished_ns, + latencia_total_ms=tts_total_ms, + latencia_tffb_ms=tts_tffb_ms, + duracao_audio_ms=tts_audio_duration_ms, + tts_max_gap_ms=tts_max_gap_ms, + tts_max_underrun_0ms=tts_max_underrun_0ms, + tts_underflow_count=tts_underflow_count, + tts_avg_underflow_ms=tts_avg_underflow_ms, + interrupcao=True if tts_interrupted else None, + erro_msg=erro_msg, + http_cod_status=200, + http_cod_desc="OK", + ) + log_flow_event( + self._call_logger, + "tts_done", + stage=stage, + message_id=message_id, + vocalized=True, + interrupted=bool(getattr(handle, "interrupted", False)), + interruption=tts_interrupted, + duration_ms=tts_total_ms, + provider_duration_ms=tts_metric_duration_ms, + metrics_source=tts_metrics_source, + tffb_ms=tts_tffb_ms, + audio_ms=tts_audio_duration_ms, + max_gap_ms=tts_max_gap_ms, + max_underrun_0ms=tts_max_underrun_0ms, + underflow_count=tts_underflow_count, + avg_underflow_ms=tts_avg_underflow_ms, + regenerated=False, + ) + self.log_room_input_snapshot(reason=f"tts_done:{stage}") + if self._timeline is not None: + self._timeline.emit( + "tts_stage_playout_done", + stage=stage, + interrupted=bool(getattr(handle, "interrupted", False)), + ) + await self._publish_debug_event( + "agent.speech.finished", + stage=stage, + message_id=message_id, + text=text, + interrupted=bool(getattr(handle, "interrupted", False)), + duration_ms=tts_total_ms, + ttfb_ms=tts_tffb_ms, + audio_duration_ms=tts_audio_duration_ms, + ) + tts_playout_completed = True + except Exception as exc: + tts_finished_ns = time.time_ns() + await self._resolve_tts_tffb_ms(handle) + ( + tts_total_ms, + tts_tffb_ms, + tts_audio_duration_ms, + tts_metric_duration_ms, + tts_metrics_source, + ) = self._resolve_tts_structured_metrics(handle) + http_cod_status, http_cod_desc = _http_status_details_from_exception(exc) + tts_interrupted = _tts_was_interrupted() + tts_error_message_id = self._tts_error_message_id( + message_id=message_id + ) + log_structured_event( + self._call_logger, + self._structured_log_context, + tipo_evento=EVENT_ENVIO_MSG, + message_id=tts_error_message_id, + inicio_ns=tts_started_ns, + fim_ns=tts_finished_ns, + latencia_total_ms=tts_total_ms, + latencia_tffb_ms=tts_tffb_ms, + duracao_audio_ms=tts_audio_duration_ms, + interrupcao=True if tts_interrupted else None, + erro_msg="Falha TTS", + erro_detalhe=f"{type(exc).__name__}: {exc}", + http_cod_status=http_cod_status, + http_cod_desc=http_cod_desc, + ) + log_flow_event( + self._call_logger, + "tts_error", + stage=stage, + message_id=tts_error_message_id, + source_message_id=message_id, + duration_ms=tts_total_ms, + provider_duration_ms=tts_metric_duration_ms, + metrics_source=tts_metrics_source, + interruption=tts_interrupted, + tffb_ms=tts_tffb_ms, + audio_ms=tts_audio_duration_ms, + error=type(exc).__name__, + ) + self._call_logger.exception( + "BOT_PLAYOUT_ERROR | stage=%s | duration_ms=%s | tffb_ms=%s", + stage, + tts_total_ms, + tts_tffb_ms if tts_tffb_ms is not None else "-", + ) + + user_spoke_during = self.state.user_final_seq != seq_at_start + handle_interrupted = bool(getattr(handle, "interrupted", False)) + if playout_result_callback is not None: + playout_result_callback( + interrupted=handle_interrupted, + user_spoke_during=user_spoke_during, + ) + + self._call_logger.info( + "BOT_SAY_RESULT | stage=%s | interrupted=%s | user_spoke_during=%s", + stage, + handle_interrupted, + user_spoke_during, + ) + if self._timeline is not None: + self._timeline.emit( + "tts_stage_result", + stage=stage, + interrupted=handle_interrupted, + user_spoke_during=user_spoke_during, + ) + + if ( + tts_playout_completed + and not handle_interrupted + and not user_spoke_during + and str(stage or "").strip().upper() == "PRESENTATION" + and not allow_int + ): + post_playout_drop_context = self._protected_speech_drop_context_from_values( + reason="protected_speech_post_playout_grace", + stage=stage, + allow_interruptions=allow_int, + drop_user_input_while_speaking=True, + agent_message_type=agent_message_type, + agent_result_type=agent_result_type, + ) + if post_playout_drop_context is not None: + self._arm_post_playout_drop_window(post_playout_drop_context) + + if user_spoke_during or handle_interrupted: + if self._config.final_grace_s > 0: + await asyncio.sleep(self._config.final_grace_s) + spoken = self.execute_now(ExtractSpokenText(handle)) + skipped = not bool((spoken or "").strip()) + await self.execute( + SetPendingInterrupt( + listened_text=spoken, + skipped=skipped, + speech_id=speech_id, + ) + ) + + self._call_logger.info( + "BOT_INTERRUPTED_AFTER_START | stage=%s | skipped=%s | speech_id=%s | spoken=%r", + stage, + skipped, + speech_id or "-", + spoken, + ) + log_flow_event( + self._call_logger, + "interrupt_marked", + reason="tts_interrupted_after_start", + stage=stage, + skipped=skipped, + text=spoken, + ) + if skipped: + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="empty_spoken_text_after_tts", + stage=stage, + skipped=True, + ) + if self._timeline is not None: + self._timeline.emit( + "tts_stage_interrupted", + stage=stage, + skipped=skipped, + spoken=spoken, + speech_id=speech_id, + ) + except Exception as exc: + if tts_started_ns is None: + tts_started_ns = time.time_ns() + tts_failed_ns = time.time_ns() + await self._resolve_tts_tffb_ms(handle) + ( + tts_total_ms, + tts_tffb_ms, + tts_audio_duration_ms, + tts_metric_duration_ms, + tts_metrics_source, + ) = self._resolve_tts_structured_metrics(handle) + http_cod_status, http_cod_desc = _http_status_details_from_exception(exc) + tts_interrupted = _tts_was_interrupted() + tts_error_message_id = self._tts_error_message_id( + message_id=message_id + ) + log_structured_event( + self._call_logger, + self._structured_log_context, + tipo_evento=EVENT_ENVIO_MSG, + message_id=tts_error_message_id, + inicio_ns=tts_started_ns, + fim_ns=tts_failed_ns, + latencia_total_ms=tts_total_ms, + latencia_tffb_ms=tts_tffb_ms, + duracao_audio_ms=tts_audio_duration_ms, + interrupcao=True if tts_interrupted else None, + erro_msg="Falha TTS", + erro_detalhe=f"{type(exc).__name__}: {exc}", + http_cod_status=http_cod_status, + http_cod_desc=http_cod_desc, + ) + log_flow_event( + self._call_logger, + "tts_error", + stage=stage, + message_id=tts_error_message_id, + source_message_id=message_id, + duration_ms=tts_total_ms, + provider_duration_ms=tts_metric_duration_ms, + metrics_source=tts_metrics_source, + interruption=tts_interrupted, + tffb_ms=tts_tffb_ms, + audio_ms=tts_audio_duration_ms, + error=type(exc).__name__, + ) + self._call_logger.exception( + "BOT_SAY_ERROR | stage=%s | duration_ms=%s | tffb_ms=%s", + stage, + tts_total_ms, + tts_tffb_ms if tts_tffb_ms is not None else "-", + ) + raise + finally: + self.state.current_speech.handle = None + self.state.current_speech.stage = "" + self.state.current_speech.text = "" + self.state.current_speech.speech_id = "" + self.state.current_speech.allow_interruptions = False + self.state.current_speech.rearm_interruptions_on_next_user_speech = False + self.state.current_speech.drop_user_input_while_speaking = False + self.state.current_speech.agent_message_type = "" + self.state.current_speech.agent_result_type = "" + self.state.current_speech.interruption_requested = False + self.state.current_speech.interruption_task = None + self.speaking.clear() + self.reset_gap_state() + if self._empty_stt_recovery_pending: + self._recover_after_empty_stt_final() + + if ( + schedule_idle_after + and self.state.user_final_seq == seq_at_start + and (stage or "").upper() != "DONE" + ): + if ( + self._idle_policy.should_arm_close_after_nudge( + stage=stage, + nudge_count=self.state.idle_nudge_count, + ) + ): + self._call_logger.info( + "IDLE_NUDGE_MAX_REACHED_AFTER_SAY | tries=%s | will_close_in_s=%.1f", + self.state.idle_nudge_count, + self._idle_policy.close_delay_s, + ) + self.arm_idle_close_timer(reason="max_reached_after_nudge") + return True + + self.arm_idle_timer(reason=f"bot_end:{(stage or '').upper()}") + + if not allow_int and not drop_user_input_while_speaking: + async with self.pending_lock: + pending_turn = self.state.pending_user_final + self.state.pending_user_final = None + if pending_turn: + self._create_task_logged( + self.run_pipeline( + pending_turn.transcription, + pending_turn.text, + user_seq=pending_turn.seq, + message_id=pending_turn.message_id, + ), + name=f"run_pipeline_pending_{pending_turn.seq}", + ) + + return True + + async def finalize_when_room_empty(self) -> None: + while not self.finalized.is_set(): + await asyncio.sleep(0.5) + try: + if self._finalization_policy.should_finalize_room_empty( + remote_participants=len(self._ctx.room.remote_participants), + finalized=self.finalized.is_set(), + ): + await asyncio.sleep(self._finalization_policy.call_end_grace_s) + if self._finalization_policy.should_finalize_room_empty( + remote_participants=len(self._ctx.room.remote_participants), + finalized=self.finalized.is_set(), + ): + await self.finalize("room_empty") + return + except Exception: + continue + + async def run_pipeline( + self, + transcription: str, + text: str, + *, + user_seq: int, + message_id: str = "", + inflight_initial_notice_consumed: bool = False, + is_deferred_replay: bool = False, + ) -> None: + text = (text or "").strip() + transcription = transcription or "" + message_id = str(message_id or "").strip() + if not message_id: + message_id = self._consume_turn_message_id(transcription, text) + elif text: + self._clear_started_turn_message_id(message_id) + + if self.should_ignore_backchannel(text): + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="backchannel", + message_id=message_id, + user_seq=user_seq, + stage=self.state.current_stage, + speaking=self.speaking.is_set(), + text=text, + ) + if is_deferred_replay: + self._reset_deferred_interruption() + return + + await self._agent._ready.wait() + if self._agent.pipeline is None: + self._logger.exception("[pipeline] pipeline=None (ready)") + if is_deferred_replay: + self._reset_deferred_interruption() + return + + async with self._agent._run_lock: + if self.finalized.is_set() or self._finalization_policy.is_done_stage( + self.state.current_stage + ): + log_flow_event( + self._call_logger, + "pipeline_run_skipped_terminal", + user_seq=user_seq, + stage=self.state.current_stage, + ) + if is_deferred_replay: + self._reset_deferred_interruption() + return + + if user_seq != self.state.user_final_seq: + self._call_logger.info( + "RUN_PIPELINE_SKIP_STALE | user_seq=%s | current_seq=%s | text=%r", + user_seq, + self.state.user_final_seq, + text, + ) + if self._timeline is not None: + self._timeline.emit( + "pipeline_run_skipped_stale", + user_seq=user_seq, + current_seq=self.state.user_final_seq, + ) + if is_deferred_replay: + self._reset_deferred_interruption() + return + + pending_interrupt = await self._agent.consume_pending_interrupt() + if len(pending_interrupt) == 3: + listened_text, speech_id, skipped = pending_interrupt + else: + listened_text, skipped = pending_interrupt + speech_id = "" + try: + if listened_text is not None: + if is_deferred_replay: + await self.execute( + SetPipelineProcessingInterruption( + listened_text=listened_text, + skipped=skipped, + speech_id=speech_id, + ) + ) + else: + await self.execute( + SetPipelineInterruption( + interrupted=True, + listened_text=listened_text, + skipped=skipped, + speech_id=speech_id, + ) + ) + else: + await self.execute(SetPipelineInterruption(interrupted=False)) + except Exception: + self._logger.exception("PIPELINE_INTERRUPT | erro chamando set_interruption") + + try: + stt_payload = json.loads(transcription) + except Exception: + stt_payload = text + stt_payload = self._payload_with_message_id(stt_payload, message_id, text=text) + + log_flow_event( + self._call_logger, + "agent_start", + message_id=message_id, + user_seq=user_seq, + stage=self.state.current_stage, + text=text, + payload=stt_payload, + ) + if self._timeline is not None: + self._timeline.emit( + "pipeline_run_started", + user_seq=user_seq, + input_text=text, + current_stage=self.state.current_stage, + ) + if text and self._supports_inflight_backend_push(): + self._schedule_backend_push_listener( + user_seq=user_seq, + reason="run_pipeline_inflight", + inflight=True, + ) + backend_started_ns = time.time_ns() + deferred = self.state.deferred_interruption + # A janela do ciclo cobre o backend e a espera de liquidacao. Depois + # dela a resposta ja esta liberada e a fala do cliente volta a ser + # barge-in comum. + deferred.backend_in_flight = not is_deferred_replay + if deferred.backend_in_flight: + self._promote_pre_backend_vad_utterances() + try: + raw_reply = await self._execute_pipeline_with_inflight_backend_wait( + stt_payload, + user_seq=user_seq, + message_id=message_id, + initial_notice_consumed=inflight_initial_notice_consumed, + ) + except InflightBackendWaitTimedOut: + self._close_deferred_vad_window() + self._release_deferred_cycle(is_deferred_replay=is_deferred_replay) + return + except Exception: + self._close_deferred_vad_window() + backend_duration_ms = event_latency_ms(backend_started_ns) + log_backend_error = getattr(self._call_logger, "error", self._call_logger.exception) + log_backend_error( + "AGENT_BACKEND_ERROR | user_seq=%s | stage=%s | duration_ms=%s", + user_seq, + self.state.current_stage, + backend_duration_ms, + ) + log_flow_event( + self._call_logger, + "agent_error", + user_seq=user_seq, + stage=self.state.current_stage, + duration_ms=backend_duration_ms, + ) + self._logger.exception("[pipeline] backend run failed") + if self._timeline is not None: + self._timeline.emit( + "pipeline_run_failed", + user_seq=user_seq, + resource="agent_backend", + ) + self._release_deferred_cycle(is_deferred_replay=is_deferred_replay) + await self.terminate_with_resource_stop( + resource="agent_backend", + source="run_pipeline", + ) + return + + try: + backend_duration_ms = event_latency_ms(backend_started_ns) + reply = self._coerce_backend_reply( + raw_reply, + fallback_stage=self.state.current_stage, + ) + reply = self._backend_reply_with_message_id(reply, message_id) + if not is_deferred_replay: + await self._wait_deferred_interruption_settlement(user_seq=user_seq) + finally: + self._close_deferred_vad_window() + + terminal_reply = self._terminal_stop_from_reply(reply) is not None or ( + reply.done or self._finalization_policy.is_done_stage(reply.stage) + ) + if user_seq != self.state.user_final_seq and not terminal_reply: + self._call_logger.info( + "RUN_PIPELINE_DROP_AFTER_RUN | user_seq=%s | current_seq=%s | stage=%s", + user_seq, + self.state.user_final_seq, + reply.stage, + ) + if self._timeline is not None: + self._timeline.emit( + "pipeline_run_dropped_after_run", + user_seq=user_seq, + current_seq=self.state.user_final_seq, + stage=reply.stage, + ) + if not is_deferred_replay and await self._dispatch_deferred_interruption(reply): + log_flow_event( + self._call_logger, + "pipeline_to_tts_skipped", + stage=reply.stage, + reason="deferred_interruption", + speech_id=self._speech_id_from_reply(reply), + ) + self.reset_gap_state() + self._release_deferred_cycle(is_deferred_replay=is_deferred_replay) + return + + self.state.current_stage = reply.stage + self._update_backend_processing_drop_from_reply(reply, source="run_pipeline") + out = reply.text + log_flow_event( + self._call_logger, + "agent_done", + message_id=message_id, + user_seq=user_seq, + stage=self.state.current_stage, + done=reply.done, + duration_ms=backend_duration_ms, + text=out, + result=reply.export_payload, + ) + if self._timeline is not None: + self._timeline.emit( + "pipeline_run_completed", + user_seq=user_seq, + stage=self.state.current_stage, + output_len=len(out), + ) + + self._emit_debug_event( + "agent.completed", + user_seq=user_seq, + message_id=message_id, + stage=self.state.current_stage, + duration_ms=backend_duration_ms, + done=reply.done, + ) + + if not out: + self._log_backend_reply_error_event( + reply, + source="run_pipeline", + message_id=message_id, + ) + + self._release_deferred_cycle(is_deferred_replay=is_deferred_replay) + if out: + self.state.gap.active = True + self.state.gap.cancelled = False + self.state.gap.guard_seq = user_seq + self.state.gap.stage = self.state.current_stage + self.state.gap.text = out + self.state.gap.speech_id = self._speech_id_from_reply(reply) + + if ( + self.state.user_final_seq != self.state.gap.guard_seq + and not terminal_reply + ): + await self.execute( + SetPendingInterrupt( + listened_text="", + skipped=True, + speech_id=self.state.gap.speech_id, + ) + ) + self.state.gap.cancelled = True + self._call_logger.info( + "SKIPPED_BETWEEN_RUN_AND_TTS | stage=%s | reason=user_spoke_during_run", + self.state.current_stage, + ) + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="user_spoke_during_agent_run", + user_seq=user_seq, + stage=self.state.current_stage, + skipped=True, + ) + if self._timeline is not None: + self._timeline.emit( + "pipeline_to_tts_skipped", + stage=self.state.current_stage, + reason="user_spoke_during_run", + ) + self.reset_gap_state() + self._release_deferred_cycle(is_deferred_replay=is_deferred_replay) + return + + if not await self._wait_for_stable_user_silence_before_terminal_reply( + reply, + user_seq=user_seq, + source="run_pipeline", + ): + self.reset_gap_state() + self._release_deferred_cycle(is_deferred_replay=is_deferred_replay) + return + + if terminal_reply: + # Um STT que chega na janela entre backend e TTS marca o gap + # como cancelado. Para respostas terminais esse marcador serve + # apenas para esperar o usuario terminar; nao pode suprimir a + # confirmacao final. + self.state.gap.cancelled = False + self.state.gap.guard_seq = self.state.user_final_seq + + try: + await self._speak_backend_reply( + reply, + source="run_pipeline", + add_to_chat_ctx=True, + initial_turn=user_seq == 0, + ) + except Exception: + self._logger.exception("[tts] failed speaking backend reply") + if self._timeline is not None: + self._timeline.emit( + "tts_stage_failed", + stage=self.state.current_stage, + source="run_pipeline", + ) + self._release_deferred_cycle(is_deferred_replay=is_deferred_replay) + await self.terminate_with_resource_stop( + resource="tts", + source="run_pipeline_tts", + ) + return + self.reset_gap_state() + self._release_deferred_cycle(is_deferred_replay=is_deferred_replay) + self.arm_agent_wait_timeout_from_reply( + reply, + user_seq=user_seq, + source="run_pipeline", + ) + + terminal_stop = self._terminal_stop_from_reply(reply) + if terminal_stop is not None: + await self._finish_with_agent_final_stop( + terminal_stop, + source="run_pipeline", + ) + return + + if reply.done or self._finalization_policy.is_done_stage(reply.stage): + self._create_task_logged( + self.notify_bridge_done("stage_done"), + name="bridge_done", + ) + self._create_task_logged( + self.finalize("stage_done"), + name="finalize", + ) + elif user_seq == self.state.user_final_seq: + self._schedule_backend_push_listener( + user_seq=user_seq, + reason="run_pipeline_completed", + ) + + async def stash_pending_user_turn( + self, + *, + seq: int, + transcription: str, + text: str, + message_id: str = "", + ) -> None: + async with self.pending_lock: + self.state.pending_user_final = PendingUserTurn( + seq=seq, + transcription=transcription, + text=text, + message_id=message_id, + ) + + def _consume_deferred_empty_final_latch(self) -> bool: + latched_at = self._deferred_empty_final_latched_at + if not latched_at: + return False + self._deferred_empty_final_latched_at = 0.0 + # A marca vale so para o evento imediatamente seguinte: se o LiveKit + # filtrou a copia, ela nao pode sobrar e engolir o final vazio de outra + # fala mais adiante. + return (time.monotonic() - latched_at) <= EMPTY_STT_FINAL_DEDUP_WINDOW_S + + def _handle_deferred_stt_final(self, *, transcription: str, text: str) -> bool: + if not self.state.deferred_interruption.backend_in_flight: + return False + utterance = self._consume_vad_for_stt_final() + if utterance is None: + return False + + message_id = self._consume_turn_message_id(transcription, text) if text else "" + if not message_id: + self._clear_started_turn_message_id() + + if not utterance.is_long: + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="deferred_interruption_too_short", + message_id=message_id, + stage=self.state.current_stage, + audio_ms=utterance.duration_ms, + minimum_audio_ms=self._config.deferred_interruption_min_audio_ms, + text=text, + ) + return True + + if not text: + log_flow_event( + self._call_logger, + "deferred_interruption_stt_empty", + stage=self.state.current_stage, + audio_ms=utterance.duration_ms, + ) + return True + + self.cancel_idle_timer("deferred_interruption") + self.cancel_idle_close("deferred_interruption") + self.cancel_agent_wait_timeout("deferred_interruption") + self._cancel_backend_push_listener(reason="deferred_interruption") + self.state.user_final_seq += 1 + event_seq = self.state.user_final_seq + self.state.last_user_final_at = time.monotonic() + deferred = self.state.deferred_interruption + deferred.long_turns.append( + PendingUserTurn( + seq=event_seq, + transcription=transcription, + text=text, + message_id=message_id, + ) + ) + log_flow_event( + self._call_logger, + "deferred_interruption_accumulated", + message_id=message_id, + user_seq=event_seq, + stage=self.state.current_stage, + audio_ms=utterance.duration_ms, + accumulated_turns=len(deferred.long_turns), + text=text, + ) + return True + + def on_user_input_transcribed(self, ev: UserInputTranscribedEvent) -> None: + transcription = ev.transcript or "" + text = (self._extract_text_from_transcript(transcription) or "").strip() + + if not ev.is_final: + return + + deferred = self.state.deferred_interruption + if deferred.replay_in_flight: + self._consume_vad_for_stt_final() + dropped_message_id = self._consume_turn_message_id(transcription, text) if text else "" + if not dropped_message_id: + self._clear_started_turn_message_id() + log_flow_event( + self._call_logger, + "deferred_replay_stt_discarded", + message_id=dropped_message_id, + text=text, + ) + self._clear_structured_interruption() + return + + if deferred.backend_in_flight and not self._config.deferred_interruption_enabled: + dropped_message_id = ( + self._consume_turn_message_id(transcription, text) if text else "" + ) + if not dropped_message_id: + self._clear_started_turn_message_id() + self._log_user_input_dropped( + context={ + "reason": "backend_processing_feedback_compat", + "stage": self.state.current_stage, + "agent_message_type": "feedback", + "agent_result_type": "feedback", + }, + transcription=transcription, + text=text, + ) + self._clear_structured_interruption() + return + + # Finais normais consomem o VAD pré-R1 correspondente. Se R1 já + # estiver em voo, _handle_deferred_stt_final faz esse consumo e + # acumula o turno longo para o replay R2. + if not deferred.backend_in_flight: + self._consume_vad_for_stt_final() + + if not text and self._consume_deferred_empty_final_latch(): + # O provider ja contabilizou este final vazio; este evento e a copia + # que o LiveKit repassa. + self._clear_structured_interruption() + return + + self._restore_agent_wait_timeout_retry_vad_threshold(reason="user_final") + drop_context = self._drop_context_for_user_final() + if drop_context is not None: + # A fala protegida tem prioridade sobre o ciclo diferido, mas a fala + # VAD correspondente precisa sair da fila junto: senao o proximo + # final consumiria a fala errada. + if deferred.backend_in_flight: + self._consume_vad_for_stt_final() + dropped_message_id = ( + self._consume_turn_message_id(transcription, text) + if text + else "" + ) + if not dropped_message_id: + self._clear_started_turn_message_id() + self._log_user_input_dropped( + context=drop_context, + transcription=transcription, + text=text, + ) + self._clear_structured_interruption() + return + + if self._handle_deferred_stt_final(transcription=transcription, text=text): + self._clear_structured_interruption() + return + + if not text: + self._clear_started_turn_message_id() + self.handle_empty_stt_final( + transcript_len=len(transcription), + source="user_input_transcribed", + ) + self._clear_structured_interruption() + return + + if deferred.backend_in_flight: + # Dentro da janela do ciclo todo final precisa de uma fala VAD + # correspondente; sem ela a associacao por ordem ja esta + # dessincronizada e reenviar este texto misturaria os turnos. Fora da + # janela o final segue o fluxo normal de interrupcao. + message_id = self._consume_turn_message_id(transcription, text) + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="deferred_interruption_missing_vad", + message_id=message_id, + stage=self.state.current_stage, + text=text, + ) + self._clear_structured_interruption() + return + + self.cancel_idle_timer("user_final") + message_id = self._consume_turn_message_id(transcription, text) + self.cancel_idle_close("user_final") + self.cancel_agent_wait_timeout("user_final") + self._clear_agent_wait_timeout_memory() + self._cancel_backend_push_listener(reason="user_final") + self.state.idle_nudge_count = 0 + + self.state.user_final_seq += 1 + event_seq = self.state.user_final_seq + self.state.last_user_final_at = time.monotonic() + was_speaking = self.speaking.is_set() + keep_pre_backend_notice = ( + was_speaking and self._pre_backend_wait_notice_is_current_speech() + ) + interruptible = ( + self._current_speech_allows_interruption() + if was_speaking + else self._interrupt_policy.allow_stage(self.state.current_stage) + ) + if keep_pre_backend_notice: + interruptible = False + log_flow_event( + self._call_logger, + "stt_final", + message_id=message_id, + user_seq=event_seq, + stage=self.state.current_stage, + speaking=was_speaking, + interruptible=interruptible, + text=text, + ) + if self._timeline is not None: + self._timeline.emit( + "user_transcript_final", + user_seq=event_seq, + message_id=message_id, + text=text, + speaking=self.speaking.is_set(), + current_stage=self.state.current_stage, + ) + stt_duration_ms = ( + round((time.monotonic() - self._last_vad_speech_end_at) * 1000) + if self._last_vad_speech_end_at > 0 + else None + ) + self._emit_debug_event( + "stt.completed", + user_seq=event_seq, + message_id=message_id, + duration_ms=stt_duration_ms, + audio_duration_ms=self._last_vad_speech_duration_ms, + text=text, + ) + + interrupt_tasks: list[asyncio.Task[Any]] = [] + if keep_pre_backend_notice: + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="pre_backend_wait_notice_same_turn", + action="keep_speaking", + message_id=message_id, + user_seq=event_seq, + stage=self.state.current_speech.stage, + text=text, + ) + + if self.state.gap.active and not self.state.gap.cancelled: + self.state.gap.cancelled = True + interrupt_tasks.append(self._create_task_logged( + self.execute( + SetPendingInterrupt( + listened_text="", + skipped=True, + speech_id=self.state.gap.speech_id, + ) + ), + name="pending_interrupt_gap", + )) + self._call_logger.info( + "GAP_INTERRUPTION | stage=%s | skipped=True | speech_id=%s | reason=user_spoke_before_tts", + self.state.gap.stage or self.state.current_stage, + self.state.gap.speech_id or "-", + ) + log_flow_event( + self._call_logger, + "interrupt_marked", + reason="user_spoke_before_tts", + message_id=message_id, + user_seq=event_seq, + stage=self.state.gap.stage or self.state.current_stage, + skipped=True, + text=text, + ) + + if was_speaking and interruptible: + log_flow_event( + self._call_logger, + "interrupt_detected", + reason="user_final_while_speaking", + message_id=message_id, + user_seq=event_seq, + stage=self.state.current_stage, + text=text, + ) + stop_task = self.request_current_speech_interruption(reason="user_final_while_speaking") + if stop_task is not None: + interrupt_tasks.append(stop_task) + interrupt_tasks.append( + self.mark_interrupt_from_handle_if_any(reason="user_final_while_speaking") + ) + elif self._last_interrupt_task is not None and not self._last_interrupt_task.done(): + if not keep_pre_backend_notice: + interrupt_tasks.append(self._last_interrupt_task) + + if was_speaking and not interruptible and not keep_pre_backend_notice: + log_flow_event( + self._call_logger, + "interrupt_ignored", + reason="stage_not_interruptible", + action="stash_pending", + message_id=message_id, + user_seq=event_seq, + stage=self.state.current_stage, + text=text, + ) + self._create_task_logged( + self.stash_pending_user_turn( + seq=event_seq, + transcription=transcription, + text=text, + message_id=message_id, + ), + name=f"stash_pending_final_{event_seq}", + ) + self._clear_structured_interruption() + return + + async def _run_after_interrupt_mark() -> None: + if interrupt_tasks: + try: + seen: set[asyncio.Task[Any]] = set() + unique_tasks: list[asyncio.Task[Any]] = [] + for task in interrupt_tasks: + if task not in seen: + seen.add(task) + unique_tasks.append(task) + await asyncio.gather(*unique_tasks) + except Exception: + self._logger.exception("PIPELINE_INTERRUPT | erro marcando interrupcao") + await self.run_pipeline(transcription, text, user_seq=event_seq, message_id=message_id) + + self._create_task_logged( + _run_after_interrupt_mark(), + name=f"run_pipeline_{event_seq}", + ) + self._clear_structured_interruption() + + async def on_shutdown(self) -> None: + await self.finalize("shutdown_callback") + + def register_callbacks(self) -> None: + @self._session.on("user_input_transcribed") + def _on_user_input_transcribed(ev: UserInputTranscribedEvent) -> None: + self.on_user_input_transcribed(ev) + + @self._session.on("error") + def _on_session_error(ev: Any) -> None: + error = getattr(ev, "error", None) + self._record_retryable_tts_session_error(error) + if self._timeline is not None: + self._timeline.emit( + "session_error", + error_type=str(getattr(error, "type", "") or ""), + recoverable=bool(getattr(error, "recoverable", False)), + ) + + @self._session.on("close") + def _on_session_close(ev: Any) -> None: + error = getattr(ev, "error", None) + raw_reason = getattr(ev, "reason", "") + reason = str(getattr(raw_reason, "value", raw_reason) or "") + if self._timeline is not None: + self._timeline.emit( + "session_closed", + reason=reason, + error_type=str(getattr(error, "type", "") or ""), + ) + if reason != "error" or error is None: + return + + resource = self._resource_from_session_error(error) + self._create_task_logged( + self.terminate_with_resource_stop( + resource=resource, + source="session_close", + ), + name=f"terminal_stop_session_close_{resource}", + ) + + @self._session.on("metrics_collected") + def _on_metrics_collected(ev: Any) -> None: + self._record_tts_metric(ev) + + @self._session.on("agent_state_changed") + def _on_agent_state_changed(ev: Any) -> None: + new_state = str(getattr(ev, "new_state", "") or "").lower() + state_name = new_state.rsplit(".", 1)[-1] + self._agent_state_name = state_name or new_state + if state_name == "speaking": + self._agent_speaking_seq += 1 + self._agent_speaking_at = time.monotonic() + self._agent_speaking_event.set() + if self._timeline is not None: + self._timeline.emit( + "agent_state_changed", + old_state=str(getattr(ev, "old_state", "") or ""), + new_state=str(getattr(ev, "new_state", "") or ""), + ) + + @self._session.on("user_state_changed") + def _on_user_state_changed(ev) -> None: + new_state = str(getattr(ev, "new_state", "") or "").lower() + state_name = new_state.rsplit(".", 1)[-1] + self._user_state_name = state_name or new_state + if state_name == "speaking": + self._user_not_speaking.clear() + else: + self._user_not_speaking.set() + self._call_logger.info( + "USER_STATE | old=%s | new=%s", + getattr(ev, "old_state", ""), + getattr(ev, "new_state", ""), + ) + if self._timeline is not None: + self._timeline.emit( + "user_state_changed", + old_state=str(getattr(ev, "old_state", "") or ""), + new_state=str(getattr(ev, "new_state", "") or ""), + ) + if state_name == "speaking": + if self.state.deferred_interruption.replay_in_flight: + log_flow_event( + self._call_logger, + "deferred_replay_user_speech_discarded", + source="user_state_changed", + ) + return + if ( + self.state.deferred_interruption.backend_in_flight + and not self._config.deferred_interruption_enabled + ): + log_flow_event( + self._call_logger, + "backend_processing_user_speech_discarded", + source="user_state_changed", + mode="feedback_compat", + ) + return + # O STT final pode chegar enquanto o conforto iniciado pelo + # mesmo turno ainda toca. So uma nova transicao real do VAD + # para speaking pode retirar essa protecao. + self._rearm_current_speech_interruptions_on_new_user_speech() + if self._mark_drop_next_user_final_from_current_speech( + reason="user_state_speaking_protected_speech" + ): + return + if self._mark_drop_next_user_final_from_post_playout_window( + reason="user_state_speaking_protected_post_playout_grace" + ): + return + self._clear_pending_user_input_drop() + if self._inflight_backend_wait_notice_is_speaking(): + log_flow_event( + self._call_logger, + "interrupt_detected", + reason="user_state_speaking", + action="interrupt_inflight_backend_wait_notice", + stage=self.state.current_speech.stage, + ) + self.note_user_speech_activity(reason="user_state_speaking") + self.request_current_speech_interruption( + reason="user_state_speaking_inflight_backend_wait", + force=True, + ) + return + self._cancel_pre_backend_wait_notice(reason="user_state_speaking") + self.note_user_speech_activity(reason="user_state_speaking") + self.request_current_speech_interruption(reason="user_state_speaking") + + @self._session.on("overlapping_speech") + def _on_overlapping_speech(ev) -> None: + is_interruption = bool(getattr(ev, "is_interruption", False)) + self._call_logger.info( + "OVERLAPPING_SPEECH | interruption=%s | delay_s=%.3f", + is_interruption, + float(getattr(ev, "detection_delay", 0.0) or 0.0), + ) + if self._timeline is not None: + self._timeline.emit( + "overlapping_speech", + is_interruption=is_interruption, + detection_delay_ms=round(float(getattr(ev, "detection_delay", 0.0) or 0.0) * 1000), + ) + if self.state.deferred_interruption.replay_in_flight: + log_flow_event( + self._call_logger, + "deferred_replay_user_speech_discarded", + source="overlapping_speech", + ) + return + if ( + self.state.deferred_interruption.backend_in_flight + and not self._config.deferred_interruption_enabled + ): + log_flow_event( + self._call_logger, + "backend_processing_user_speech_discarded", + source="overlapping_speech", + mode="feedback_compat", + ) + return + if is_interruption: + if self._mark_drop_next_user_final_from_current_speech( + reason="overlapping_speech_protected_speech" + ): + return + self.request_current_speech_interruption(reason="overlapping_speech") + + room_on = getattr(self._ctx.room, "on", None) + if callable(room_on): + @room_on("participant_connected") + def _on_room_participant_connected(participant: Any) -> None: + log_flow_event( + self._call_logger, + "agent_room_participant_connected", + participant=str(getattr(participant, "identity", "") or ""), + kind=str(getattr(participant, "kind", "") or ""), + bridge_identity=self._bridge_identity or "", + ) + self.log_room_input_snapshot(reason="participant_connected") + + @room_on("participant_disconnected") + def _on_room_participant_disconnected(participant: Any) -> None: + log_flow_event( + self._call_logger, + "agent_room_participant_disconnected", + participant=str(getattr(participant, "identity", "") or ""), + reason=str(getattr(participant, "disconnect_reason", "") or ""), + bridge_identity=self._bridge_identity or "", + ) + self.log_room_input_snapshot(reason="participant_disconnected") + + @room_on("track_published") + def _on_room_track_published(publication: Any, participant: Any) -> None: + log_flow_event( + self._call_logger, + "agent_room_track_published", + participant=str(getattr(participant, "identity", "") or ""), + bridge_identity=self._bridge_identity or "", + **self._track_publication_fields(publication), + ) + self.log_room_input_snapshot(reason="track_published") + + @room_on("track_subscribed") + def _on_room_track_subscribed(track: Any, publication: Any, participant: Any) -> None: + fields = self._track_publication_fields(publication) + fields["track_sid"] = fields.get("track_sid") or str(getattr(track, "sid", "") or "") + log_flow_event( + self._call_logger, + "agent_room_track_subscribed", + participant=str(getattr(participant, "identity", "") or ""), + bridge_identity=self._bridge_identity or "", + **fields, + ) + self.log_room_input_snapshot(reason="track_subscribed") + + @room_on("track_unsubscribed") + def _on_room_track_unsubscribed(track: Any, publication: Any, participant: Any) -> None: + fields = self._track_publication_fields(publication) + fields["track_sid"] = fields.get("track_sid") or str(getattr(track, "sid", "") or "") + log_flow_event( + self._call_logger, + "agent_room_track_unsubscribed", + participant=str(getattr(participant, "identity", "") or ""), + bridge_identity=self._bridge_identity or "", + **fields, + ) + self.log_room_input_snapshot(reason="track_unsubscribed") + + @room_on("track_muted") + def _on_room_track_muted(participant: Any, publication: Any) -> None: + log_flow_event( + self._call_logger, + "agent_room_track_muted", + participant=str(getattr(participant, "identity", "") or ""), + bridge_identity=self._bridge_identity or "", + **self._track_publication_fields(publication), + ) + self.log_room_input_snapshot(reason="track_muted") + + @room_on("track_unmuted") + def _on_room_track_unmuted(participant: Any, publication: Any) -> None: + log_flow_event( + self._call_logger, + "agent_room_track_unmuted", + participant=str(getattr(participant, "identity", "") or ""), + bridge_identity=self._bridge_identity or "", + **self._track_publication_fields(publication), + ) + self.log_room_input_snapshot(reason="track_unmuted") + + @room_on("data_received") + def _on_room_data_received(data_packet: Any) -> None: + if getattr(data_packet, "topic", "") != "bridge.control": + return + + message = self._decode_data_packet(data_packet) + if message is None: + return + + control_type = str(message.get("type") or "").strip() + if self._timeline is not None: + self._timeline.emit("bridge_control_received", control_type=control_type) + + if control_type == "client_audio_enabled": + self.request_initial_agent_turn(reason="client_audio_enabled") + + self._ctx.add_shutdown_callback(self.on_shutdown) + + async def run(self) -> None: + self.register_callbacks() + self._create_task_logged( + self.finalize_when_room_empty(), + name="finalize_when_room_empty", + ) + + if self._agent_starts_conversation: + # Precisa vir antes do StartSession: o RoomIO respeita o estado do + # input ao anexar a track, entao o audio ja nasce descartado. + self.hold_user_audio_input(reason="room_setup") + self._arm_user_audio_input_gate_timeout() + + room_opts_kwargs = dict( + audio_input=room_io.AudioInputOptions( + sample_rate=16000, + num_channels=1, + frame_size_ms=20, + pre_connect_audio=False, + ), + audio_output=room_io.AudioOutputOptions( + sample_rate=16000, + num_channels=1, + ), + ) + + if self._bridge_identity: + room_opts_kwargs["participant_identity"] = self._bridge_identity + + self._call_logger.debug( + "ROOM_ENTER | room=%s | participant_identity=%s", + self._ctx.room.name, + self._bridge_identity or "-", + ) + if self._timeline is not None: + self._timeline.emit( + "room_enter", + participant_identity=self._bridge_identity or "", + ) + + async def _start_session() -> None: + await self.execute( + StartSession( + room=self._ctx.room, + room_options=room_io.RoomOptions(**room_opts_kwargs), + ) + ) + + t_session = self._create_task_logged(_start_session(), name="session_start") + if self._timeline is not None: + self._timeline.emit("session_start_requested") + self.arm_idle_timer(reason="client_join", delay_s=self._idle_policy.join_delay_s) + + try: + await t_session + if self._timeline is not None: + self._timeline.emit("session_started") + except Exception as exc: + self._log_resource_error_events( + status="stop_agent_runtime_unavailable", + reason="livekit_session_start_failed", + resource="agent_runtime", + source="session_start", + exc=exc, + ) + raise + finally: + self._call_logger.debug("ROOM_EXIT | room=%s", self._ctx.room.name) + if self._timeline is not None: + self._timeline.emit("room_exit") diff --git a/src/app/livekit/runtime/command_executor.py b/src/app/livekit/runtime/command_executor.py new file mode 100644 index 0000000..9f2c859 --- /dev/null +++ b/src/app/livekit/runtime/command_executor.py @@ -0,0 +1,131 @@ +from __future__ import annotations + +from typing import Any + +from app.livekit.runtime.commands import ( + EndServiceOnce, + ExportSession, + ExtractSpokenText, + InterruptSpeech, + InjectIdleNudge, + NotifyBridgeDone, + NotifyBridgeStop, + RunPipelineInput, + SetPendingInterrupt, + SetPipelineInterruption, + SetPipelineProcessingInterruption, + StartSession, + StartSpeech, + WaitForSpeechPlayout, +) + + +class RuntimeCommandExecutor: + def __init__( + self, + *, + agent: Any, + bridge_gateway: Any, + export_service: Any, + session: Any, + speech_service: Any, + ) -> None: + self._agent = agent + self._bridge_gateway = bridge_gateway + self._export_service = export_service + self._session = session + self._speech_service = speech_service + + def execute_now(self, command: Any) -> Any: + if isinstance(command, ExtractSpokenText): + return self._speech_service.extract_spoken_text(command.handle) + + raise TypeError(f"Unsupported synchronous command: {type(command)!r}") + + async def execute(self, command: Any) -> Any: + if isinstance(command, NotifyBridgeDone): + await self._bridge_gateway.notify_stage_done(command.reason) + return None + + if isinstance(command, NotifyBridgeStop): + await self._bridge_gateway.notify_stop( + status=command.status, + reason=command.reason, + resource=command.resource, + failed_resources=command.failed_resources, + phase=command.phase, + ) + return None + + if isinstance(command, ExportSession): + await self._export_service.export_session(command.output, command.session_id) + return None + + if isinstance(command, EndServiceOnce): + return await self._agent.end_service_once() + + if isinstance(command, InjectIdleNudge): + if self._agent.pipeline is None: + raise RuntimeError("Agent backend not ready") + await self._agent.pipeline.inject_idle_nudge(command.text) + return None + + if isinstance(command, SetPendingInterrupt): + await self._agent.set_pending_interrupt( + listened_text=command.listened_text, + skipped=command.skipped, + speech_id=command.speech_id, + ) + return None + + if isinstance(command, StartSpeech): + return await self._speech_service.start( + command.text, + allow_interruptions=command.allow_interruptions, + add_to_chat_ctx=command.add_to_chat_ctx, + audio=command.audio, + ) + + if isinstance(command, WaitForSpeechPlayout): + await self._speech_service.wait_for_playout(command.handle) + return None + + if isinstance(command, InterruptSpeech): + await self._speech_service.interrupt(command.handle, force=command.force) + return None + + if isinstance(command, SetPipelineInterruption): + if self._agent.pipeline is None: + raise RuntimeError("Agent backend not ready") + await self._agent.pipeline.set_interruption( + command.interrupted, + command.listened_text, + command.skipped, + command.speech_id, + ) + return None + + if isinstance(command, SetPipelineProcessingInterruption): + if self._agent.pipeline is None: + raise RuntimeError("Agent backend not ready") + await self._agent.pipeline.set_processing_interruption( + command.listened_text, + command.skipped, + command.speech_id, + ) + return None + + if isinstance(command, RunPipelineInput): + if self._agent.pipeline is None: + raise RuntimeError("Agent backend not ready") + return await self._agent.pipeline.run(command.user_input) + + if isinstance(command, StartSession): + await self._session.start( + agent=self._agent, + room=command.room, + room_options=command.room_options, + ) + return None + + raise TypeError(f"Unsupported command: {type(command)!r}") diff --git a/src/app/livekit/runtime/commands.py b/src/app/livekit/runtime/commands.py new file mode 100644 index 0000000..b867345 --- /dev/null +++ b/src/app/livekit/runtime/commands.py @@ -0,0 +1,91 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + + +@dataclass(frozen=True, slots=True) +class NotifyBridgeDone: + reason: str = "stage_done" + + +@dataclass(frozen=True, slots=True) +class NotifyBridgeStop: + status: str + reason: str + resource: str = "" + failed_resources: tuple[str, ...] = () + phase: str = "in_session" + + +@dataclass(frozen=True, slots=True) +class ExportSession: + output: Any + session_id: str + + +@dataclass(frozen=True, slots=True) +class EndServiceOnce: + pass + + +@dataclass(frozen=True, slots=True) +class InjectIdleNudge: + text: str + + +@dataclass(frozen=True, slots=True) +class SetPendingInterrupt: + listened_text: str = "" + skipped: bool = False + speech_id: str = "" + + +@dataclass(frozen=True, slots=True) +class StartSpeech: + text: str + allow_interruptions: bool + add_to_chat_ctx: bool = True + audio: Any = None + + +@dataclass(frozen=True, slots=True) +class WaitForSpeechPlayout: + handle: Any + + +@dataclass(frozen=True, slots=True) +class InterruptSpeech: + handle: Any + force: bool = False + + +@dataclass(frozen=True, slots=True) +class ExtractSpokenText: + handle: Any + + +@dataclass(frozen=True, slots=True) +class SetPipelineInterruption: + interrupted: bool + listened_text: str = "" + skipped: bool = False + speech_id: str = "" + + +@dataclass(frozen=True, slots=True) +class SetPipelineProcessingInterruption: + listened_text: str = "" + skipped: bool = False + speech_id: str = "" + + +@dataclass(frozen=True, slots=True) +class RunPipelineInput: + user_input: Any + + +@dataclass(frozen=True, slots=True) +class StartSession: + room: Any + room_options: Any diff --git a/src/app/livekit/runtime/initial_greeting_audio_cache.py b/src/app/livekit/runtime/initial_greeting_audio_cache.py new file mode 100644 index 0000000..1d93dd3 --- /dev/null +++ b/src/app/livekit/runtime/initial_greeting_audio_cache.py @@ -0,0 +1,186 @@ +from __future__ import annotations + +import asyncio +import hashlib +import os +import tempfile +import time +import wave +from dataclasses import dataclass +from pathlib import Path +from typing import AsyncIterator + + +def _env_bool(name: str, default: bool) -> bool: + raw = (os.getenv(name, "") or "").strip().lower() + if not raw: + return default + return raw in {"1", "true", "yes", "on"} + + +def _env_int(name: str, default: int) -> int: + try: + return int(os.getenv(name, str(default))) + except ValueError: + return default + + +@dataclass(frozen=True, slots=True) +class InitialGreetingCacheKey: + agent: str + text: str + provider: str + voice: str + language: str + sample_rate: int + + @property + def digest(self) -> str: + raw = "\x1f".join((self.agent, self.text, self.provider, self.voice, self.language, str(self.sample_rate))) + return hashlib.sha256(raw.encode("utf-8")).hexdigest() + + +class InitialGreetingAudioCache: + """Read-through cache for an exact, first-turn TTS response. + + The caller is responsible for only offering first-turn text. Cache entries are + keyed by the rendered text and TTS settings, so a changed greeting never + reuses stale audio. + """ + + def __init__(self) -> None: + self.enabled = _env_bool("INITIAL_GREETING_AUDIO_CACHE_ENABLED", True) + self._max_chars = max(1, _env_int("INITIAL_GREETING_AUDIO_CACHE_MAX_CHARS", 1_000)) + self._ttl_s = max(1, _env_int("INITIAL_GREETING_AUDIO_CACHE_TTL_S", 3_600)) + self._max_entries = max(1, _env_int("INITIAL_GREETING_AUDIO_CACHE_MAX_ENTRIES", 32)) + self._max_bytes = max(44, _env_int("INITIAL_GREETING_AUDIO_CACHE_MAX_BYTES", 16 * 1024 * 1024)) + self._max_agent_variants = max(1, _env_int("INITIAL_GREETING_AUDIO_CACHE_MAX_VARIANTS_PER_AGENT", 1)) + self._disable_ttl_s = max(1, _env_int("INITIAL_GREETING_AUDIO_CACHE_DISABLE_TTL_S", 300)) + self._agent_variants: dict[str, set[str]] = {} + self._agent_disabled_until: dict[str, float] = {} + self._dir = Path(os.getenv("INITIAL_GREETING_AUDIO_CACHE_DIR", "./cache/initial-greetings")) + + def key_for( + self, + *, + agent: str, + text: str, + provider: str, + voice: str, + language: str, + sample_rate: int, + ) -> InitialGreetingCacheKey | None: + normalized = " ".join(str(text or "").split()) + agent_name = (agent or "unknown").strip().lower() + now = time.monotonic() + if self._agent_disabled_until.get(agent_name, 0.0) > now: + return None + variants = self._agent_variants.setdefault(agent_name, set()) + if normalized not in variants and len(variants) >= self._max_agent_variants: + self._agent_disabled_until[agent_name] = now + self._disable_ttl_s + return None + variants.add(normalized) + if not self.enabled or not normalized or len(normalized) > self._max_chars: + return None + return InitialGreetingCacheKey( + agent=(agent or "unknown").strip().lower(), + text=normalized, + provider=(provider or "unknown").strip().lower(), + voice=(voice or "default").strip(), + language=(language or "pt-BR").strip().lower(), + sample_rate=max(1, int(sample_rate)), + ) + + def path_for(self, key: InitialGreetingCacheKey) -> Path: + return self._dir / f"{key.digest}.wav" + + def has(self, key: InitialGreetingCacheKey) -> bool: + path = self.path_for(key) + try: + stat = path.stat() + except FileNotFoundError: + return False + if stat.st_size <= 44 or time.time() - stat.st_mtime >= self._ttl_s: + self.discard(key) + return False + os.utime(path, None) + return True + + def frames(self, key: InitialGreetingCacheKey) -> AsyncIterator[object]: + return self._materialize_frames(key) + + def discard(self, key: InitialGreetingCacheKey) -> None: + self.path_for(key).unlink(missing_ok=True) + + def _materialize_frames(self, key: InitialGreetingCacheKey) -> AsyncIterator[object]: + path = self.path_for(key) + try: + with wave.open(str(path), "rb") as wav: + channels = wav.getnchannels() + sample_width = wav.getsampwidth() + sample_rate = wav.getframerate() + if channels != 1 or sample_width != 2: + raise ValueError(f"Unsupported cached greeting WAV format: {path}") + pcm = wav.readframes(10**9) + if not pcm: + raise ValueError(f"Cached greeting WAV is empty: {path}") + except (OSError, EOFError, ValueError, wave.Error): + self.discard(key) + raise + + async def _frames() -> AsyncIterator[object]: + from livekit import rtc + + samples_per_frame = max(1, round(sample_rate * 20 / 1000)) + frame_bytes = samples_per_frame * 2 * channels + for offset in range(0, len(pcm), frame_bytes): + chunk = pcm[offset : offset + frame_bytes] + if not chunk: + continue + if len(chunk) < frame_bytes: + chunk = chunk + (b"\x00" * (frame_bytes - len(chunk))) + yield rtc.AudioFrame( + data=chunk, + sample_rate=sample_rate, + num_channels=channels, + samples_per_channel=len(chunk) // (2 * channels), + ) + await asyncio.sleep(0) + + return _frames() + + async def store_pcm(self, key: InitialGreetingCacheKey, pcm: bytes) -> None: + if not pcm or self.has(key): + return + await asyncio.to_thread(self._store_pcm_sync, key, pcm) + + def _store_pcm_sync(self, key: InitialGreetingCacheKey, pcm: bytes) -> None: + self._dir.mkdir(parents=True, exist_ok=True) + target = self.path_for(key) + if self.has(key): + return + with tempfile.NamedTemporaryFile(dir=self._dir, suffix=".wav", delete=False) as handle: + tmp_path = Path(handle.name) + try: + with wave.open(str(tmp_path), "wb") as wav: + wav.setnchannels(1) + wav.setsampwidth(2) + wav.setframerate(key.sample_rate) + wav.writeframes(pcm) + os.replace(tmp_path, target) + self._enforce_limits_sync() + finally: + tmp_path.unlink(missing_ok=True) + + def _enforce_limits_sync(self) -> None: + entries = [path for path in self._dir.glob("*.wav") if path.is_file()] + for path in entries[:]: + if time.time() - path.stat().st_mtime >= self._ttl_s: + path.unlink(missing_ok=True) + entries.remove(path) + entries.sort(key=lambda path: path.stat().st_mtime) + total_bytes = sum(path.stat().st_size for path in entries) + while entries and (len(entries) > self._max_entries or total_bytes > self._max_bytes): + evicted = entries.pop(0) + total_bytes -= evicted.stat().st_size + evicted.unlink(missing_ok=True) diff --git a/src/app/livekit/runtime/scheduler.py b/src/app/livekit/runtime/scheduler.py new file mode 100644 index 0000000..bb188e8 --- /dev/null +++ b/src/app/livekit/runtime/scheduler.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +import asyncio +from dataclasses import dataclass +from typing import Any, Callable + + +@dataclass(slots=True) +class ScheduledTimer: + token: int = 0 + task: asyncio.Task | None = None + + +class TimerScheduler: + def __init__(self, *, create_task_logged: Callable[..., asyncio.Task], logger: Any) -> None: + self._create_task_logged = create_task_logged + self._logger = logger + self._timers: dict[str, ScheduledTimer] = {} + + def cancel(self, key: str, *, reason: str, log_name: str) -> None: + timer = self._timers.get(key) + if timer is None or timer.task is None: + return + if timer.task is asyncio.current_task(): + return + if not timer.task.done(): + timer.task.cancel() + self._logger.debug("%s | reason=%s", log_name, reason) + timer.task = None + + def arm( + self, + key: str, + *, + task_name: str, + coro, + ) -> int: + timer = self._timers.setdefault(key, ScheduledTimer()) + timer.token += 1 + token = timer.token + timer.task = self._create_task_logged(coro, name=task_name) + return token + + def is_current(self, key: str, token: int) -> bool: + timer = self._timers.get(key) + if timer is None: + return False + return timer.token == token diff --git a/src/app/livekit/runtime/state.py b/src/app/livekit/runtime/state.py new file mode 100644 index 0000000..b6593cc --- /dev/null +++ b/src/app/livekit/runtime/state.py @@ -0,0 +1,99 @@ +from __future__ import annotations + +from collections import deque +from dataclasses import dataclass, field +from typing import Any, Optional + + +@dataclass(slots=True) +class PendingUserTurn: + seq: int + transcription: str + text: str + message_id: str = "" + + +@dataclass(slots=True) +class DeferredVADUtterance: + """A fala encerrada pelo VAD, aguardando seu STT final correspondente.""" + + duration_ms: int + is_long: bool + ended_at: float = 0.0 + + +@dataclass(slots=True) +class DeferredInterruptionState: + # Janela do ciclo: aberta enquanto o backend do turno original esta em voo, + # fechada assim que a resposta dele e liberada para o playout. Fora dela a + # fala do cliente e uma interrupcao comum e segue o fluxo normal. + backend_in_flight: bool = False + # Finais VAD anteriores a R1 cujo STT ainda nao terminou. + pre_backend_vad_utterances: deque[DeferredVADUtterance] = field(default_factory=deque) + pending_vad_utterances: deque[DeferredVADUtterance] = field(default_factory=deque) + pending_long_stt_finals: int = 0 + long_turns: list[PendingUserTurn] = field(default_factory=list) + special_comfort_sent: bool = False + replay_in_flight: bool = False + # Incrementado a cada reset para que uma tarefa de conforto atrasada nao + # toque no ciclo seguinte. + cycle_seq: int = 0 + + +@dataclass(slots=True) +class GapState: + active: bool = False + guard_seq: int = 0 + stage: str = "" + text: str = "" + speech_id: str = "" + cancelled: bool = False + + +@dataclass(slots=True) +class CurrentSpeechState: + handle: Optional[Any] = None + stage: str = "" + text: str = "" + speech_id: str = "" + allow_interruptions: bool = False + rearm_interruptions_on_next_user_speech: bool = False + drop_user_input_while_speaking: bool = False + agent_message_type: str = "" + agent_result_type: str = "" + interruption_requested: bool = False + interruption_task: Optional[Any] = None + +@dataclass(slots=True) +class CallState: + current_stage: str = "INTRO" + pending_user_final: Optional[PendingUserTurn] = None + user_final_seq: int = 0 + last_user_final_at: float = 0.0 + idle_nudge_count: int = 0 + gap: GapState = field(default_factory=GapState) + current_speech: CurrentSpeechState = field(default_factory=CurrentSpeechState) + deferred_interruption: DeferredInterruptionState = field(default_factory=DeferredInterruptionState) + + +@dataclass(frozen=True, slots=True) +class RuntimeConfig: + call_end_grace_s: float + final_grace_s: float + idle_nudge_delay_s: float + idle_nudge_join_delay_s: float + idle_nudge_close_delay_s: float + idle_nudge_max_tries: int + idle_nudge_end_reason: str + inflight_backend_wait_interval_s: float = 0.0 + inflight_backend_wait_timeout_s: float = 0.0 + inflight_backend_wait_max_notices: int = 0 + inflight_backend_wait_text: str = "" + inflight_backend_wait_short_audio_dir: str = "" + inflight_backend_wait_long_audio_dir: str = "" + inflight_backend_wait_long_audio_path: str = "" + pre_backend_wait_notice_fast_on_vad_pause: bool = False + deferred_interruption_enabled: bool = True + deferred_interruption_min_audio_ms: int = 1000 + deferred_interruption_stt_settle_timeout_s: float = 3.0 + deferred_interruption_user_turn_timeout_s: float = 10.0 diff --git a/src/app/livekit/runtime/wav_audio.py b/src/app/livekit/runtime/wav_audio.py new file mode 100644 index 0000000..c38ca7c --- /dev/null +++ b/src/app/livekit/runtime/wav_audio.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +import asyncio +import wave +from functools import lru_cache +from pathlib import Path +from typing import Any, AsyncIterator + +DEFAULT_TAIL_SILENCE_MS = 320 + + +@lru_cache(maxsize=16) +def _read_wav_pcm(path: str) -> tuple[bytes, int, int]: + wav_path = Path(path) + with wave.open(str(wav_path), "rb") as wav: + channels = wav.getnchannels() + sample_width = wav.getsampwidth() + sample_rate = wav.getframerate() + if channels != 1 or sample_width != 2: + raise ValueError( + f"Unsupported wait audio format for {wav_path}: " + f"expected PCM16 mono, got channels={channels}, sample_width={sample_width}" + ) + data = wav.readframes(10**9) + if not data: + raise ValueError(f"Wait audio file is empty: {wav_path}") + return data, sample_rate, channels + + +async def wav_audio_frames( + path: str, + *, + frame_duration_ms: int = 20, + tail_silence_ms: int = DEFAULT_TAIL_SILENCE_MS, +) -> AsyncIterator[Any]: + from livekit import rtc + + data, sample_rate, channels = _read_wav_pcm(path) + samples_per_frame = max(1, round(sample_rate * max(1, frame_duration_ms) / 1000)) + bytes_per_sample = 2 * channels + frame_bytes = samples_per_frame * bytes_per_sample + + def _frame(chunk: bytes) -> Any: + if len(chunk) < frame_bytes: + chunk = chunk + (b"\x00" * (frame_bytes - len(chunk))) + return rtc.AudioFrame( + data=chunk, + sample_rate=sample_rate, + num_channels=channels, + samples_per_channel=len(chunk) // bytes_per_sample, + ) + + for offset in range(0, len(data), frame_bytes): + chunk = data[offset : offset + frame_bytes] + samples_per_channel = len(chunk) // bytes_per_sample + if samples_per_channel <= 0: + continue + yield _frame(chunk) + await asyncio.sleep(0) + + tail_ms = max(0, int(tail_silence_ms)) + frame_ms = max(1, int(frame_duration_ms)) + tail_frames = (tail_ms + frame_ms - 1) // frame_ms + silence = b"\x00" * frame_bytes + for _ in range(tail_frames): + yield _frame(silence) + await asyncio.sleep(0) + + +def wav_duration_ms(path: str) -> int: + data, sample_rate, channels = _read_wav_pcm(path) + samples = len(data) // (2 * channels) + return round((samples / sample_rate) * 1000) diff --git a/src/app/livekit/vad_dynamic_threshold.py b/src/app/livekit/vad_dynamic_threshold.py new file mode 100644 index 0000000..efce1dc --- /dev/null +++ b/src/app/livekit/vad_dynamic_threshold.py @@ -0,0 +1,134 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +from app.utils.logging import log_flow_event + + +@dataclass(frozen=True, slots=True) +class DynamicVADThresholdConfig: + enabled: bool + baseline_activation_threshold: float + baseline_deactivation_threshold: float + retry_activation_threshold: float + retry_deactivation_threshold: float + + +class DynamicVADThresholdController: + def __init__( + self, + vad: Any, + *, + config: DynamicVADThresholdConfig, + call_logger: Any = None, + timeline: Any = None, + ) -> None: + self._vad = vad + self._config = config + self._call_logger = call_logger + self._timeline = timeline + self._retry_active = False + self._current_mode = "baseline" + self._current_activation_threshold = config.baseline_activation_threshold + self._current_deactivation_threshold = config.baseline_deactivation_threshold + + @property + def enabled(self) -> bool: + return bool(self._config.enabled) + + @property + def retry_active(self) -> bool: + return self._retry_active + + def current_thresholds(self) -> dict[str, Any]: + return { + "mode": self._current_mode, + "activation_threshold": self._current_activation_threshold, + "deactivation_threshold": self._current_deactivation_threshold, + } + + def activate_agent_wait_timeout_retry(self, *, attempt: int, reason: str) -> bool: + if not self.enabled or self._retry_active: + return False + applied = self._apply( + activation_threshold=self._config.retry_activation_threshold, + deactivation_threshold=self._config.retry_deactivation_threshold, + mode="agent_wait_timeout_retry", + reason=reason, + attempt=attempt, + ) + if applied: + self._retry_active = True + return applied + + def restore(self, *, reason: str) -> bool: + if not self.enabled or not self._retry_active: + return False + applied = self._apply( + activation_threshold=self._config.baseline_activation_threshold, + deactivation_threshold=self._config.baseline_deactivation_threshold, + mode="baseline", + reason=reason, + attempt=None, + ) + if applied: + self._retry_active = False + return applied + + def _apply( + self, + *, + activation_threshold: float, + deactivation_threshold: float, + mode: str, + reason: str, + attempt: int | None, + ) -> bool: + update_options = getattr(self._vad, "update_options", None) + if not callable(update_options): + self._emit( + "vad_dynamic_threshold_failed", + mode=mode, + reason=reason, + attempt=attempt, + error="update_options_unavailable", + ) + return False + + try: + update_options( + activation_threshold=activation_threshold, + deactivation_threshold=deactivation_threshold, + ) + except Exception as exc: + if self._call_logger is not None: + self._call_logger.exception("[vad] failed updating dynamic threshold") + self._emit( + "vad_dynamic_threshold_failed", + mode=mode, + reason=reason, + attempt=attempt, + error=type(exc).__name__, + ) + return False + + self._current_mode = mode + self._current_activation_threshold = activation_threshold + self._current_deactivation_threshold = deactivation_threshold + self._emit( + "vad_dynamic_threshold_updated", + mode=mode, + reason=reason, + attempt=attempt, + activation_threshold=activation_threshold, + deactivation_threshold=deactivation_threshold, + ) + return True + + def _emit(self, event: str, **fields: Any) -> None: + clean_fields = {key: value for key, value in fields.items() if value is not None} + if self._call_logger is not None: + log_flow_event(self._call_logger, event, **clean_fields) + if self._timeline is not None: + self._timeline.emit(event, **clean_fields) diff --git a/src/app/models/schemas.py b/src/app/models/schemas.py new file mode 100644 index 0000000..cc082bb --- /dev/null +++ b/src/app/models/schemas.py @@ -0,0 +1,46 @@ +from pydantic import BaseModel +from typing import Literal, Dict + + +class AgentData(BaseModel): + idFatura: str | None = None + current_invoice_number: str | None = None + currentInvoiceNumber: str | None = None + channel: str | None = None + + +class StartData(BaseModel): + agent: str + ani: str + gsm: str + routerCallKeyDay: str + routerCallKey: str + callIdGed: str + session_id: str + sessionId: str | None = None + protocolo: str | None = None + agentData: AgentData | None = None + + +class AudioFormat(BaseModel): + encoding: str + sampleRateHz: int + channels: int + +class StartMessage(BaseModel): + type: Literal["start"] + data: StartData + audioFormat: Dict[str, object] | AudioFormat | None = None + callConfig: Dict[str, object] | None = None + +class StopMessage(BaseModel): + type: Literal["stop"] + + +class TransferenciaSessionIdData(BaseModel): + session_id: str + + +class TransferenciaSessionIdMessage(BaseModel): + type: Literal["transferencia_session_id"] + data: TransferenciaSessionIdData diff --git a/src/app/models/stages.py b/src/app/models/stages.py new file mode 100644 index 0000000..0a3cfe5 --- /dev/null +++ b/src/app/models/stages.py @@ -0,0 +1,18 @@ +from enum import Enum + +class Stages(Enum): + presentation = "Faz perguntas do usuário para saber se está falando com o cliente alvo, pode receber respostas do tipo: sim, não, sou eu, espere um pouco etc." + argumentation = "Conversa com o usuário tentando vender um plano de telefone, possui algumas tentativas de vendas, pode receber respostas do tipo: tipo: sim, não, aceito comprar, não quero, não tenho interesse etc." + data_confirmation = "Confirma os dados do cliente, pode receber nomes, datas e números." + formalization = "Faz a confirmação final da venda, vocalizando um pedindo para o cliente confirmar, pode receber respostas do tipo: sim, não, aceito comprar, não quero, não tenho interesse etc." + + @classmethod + def get_value(cls, key: str = "") -> str: + """ + Retorna o value (descrição) direto a partir do nome da stage. + Ex: Stages.get_value("formalization") + """ + try: + return cls[key].value + except KeyError: + return "" diff --git a/src/app/providers/stt_config.py b/src/app/providers/stt_config.py new file mode 100644 index 0000000..bd5f30d --- /dev/null +++ b/src/app/providers/stt_config.py @@ -0,0 +1,551 @@ +override_config = { + "pre_processors": [ + { + "strategy": "selective_short_low_peak", + "config": {} + } + ], + "processor": { + "strategy": "faster_default", + "config": { + "extra_params": { + "temperature": 0, + "without_timestamps": True, + "word_timestamps": False + }, + "model": "igorcouto/sofya_sft_600h_no_spec_augment_v2-ctranslate2" + } + }, + "post_processors": [ + { + "strategy": "similarity_errors_default", + "config": { + "threshold": 90, + "errors": { + "Já, confio": "É confirmo", + "Ja, confio": "É confirmo", + "confio": "confirmo", + "Confio": "Confirmo", + "YIN": "TIM", + "a tema": "a TIM", + "a Teana": "a TIM", + "Timpratx": "TIM pra TIM", + "Atendimento e controle": "Atendimento TIM Controle", + "Assistente Batinho": "Assistente TIM", + "Bom dia, putinho": "Bom dia, TIM", + "loja ti ninha": "Loja TIM", + "débite": "débito", + "a Lagoa": "Alagoas", + "DD": "DDD", + "código de barro": "código de barras", + "cliente tia": "cliente TIM", + "Ouvidoria, tio": "Ouvidoria TIM", + "Nome complexo": "Nome completo", + "Pelo assinamento especial": "Atendimento especial", + "operadora Eixinha": "operadora TIM", + "a datinha": "a da TIM", + "prepédio": "pré-pago", + "Aprendimento do time": "Atendimento TIM", + "Dimei, Tutim": "Atendimento TIM", + "Servimento Saque": "Atendimento SAC", + "Período de Chimbo": "Atendimento TIM", + "Convidoria aqui": "Ouvidoria TIM", + "Comprindo seu nome": "Confirma seu nome", + "cliente na Aline": "cliente na linha", + "mês é atual titular": "mas é atual titular", + "Bem, Aitinho": "Atendimento TIM", + "chique": "chip", + "lá na Goa": "Alagoas", + "Ação de crédito": "cartão de crédito", + "nova recada": "nova recarga", + "Pagamento via fixo": "pagamento via pix", + "Miltin": "meu TIM", + "Paranáute": "Paramount", + "para mont": "Paramount", + "outurro": "outubro", + "operador Petinho": "operadora TIM", + "cayer": "cair", + "pás- pago": "pós-pago", + "Ligação admitada": "Ligação ilimitada", + "telefono": "telefone", + "a VULSO": "avulso", + "ligização": "ligação", + "aparelho das horas ": "aparelho da senhora", + "alugação cair": "ligação cair", + "pusei um código": "puxei um código", + "duas contras": "duas contas", + "No meu interior": "No meu anterior", + "Atim": "a TIM", + "acidente anterior": "atendente anterior", + "Te encontrou-lhe": "TIM Controle", + "Central Tema": "Central TIM", + "meu tema": "meu TIM", + "aplicativo Estufim": "aplicativo Meu TIM", + "aplicativo do meu time": "aplicativo do Meu TIM", + "aplicativo meu tinha": "aplicativo Meu TIM", + "w.tema.com.BR": "www.tim.com.br", + "notando o número": "anotando o número", + "time controle": "TIM Controle", + "TeamControle": "TIM Controle", + "TeamControleSmart": "TIM Controle Smart", + "Tema Com trol Ela": "TIM Controle", + "Tinha controle redes sociais": "TIM Controle redes sociais", + "tinha controle de redes sociais": "TIM Controle redes sociais", + "fazer uma potabilidade": "fazer uma portabilidade", + "Ablaca": "Black", + "Meu Team": "Meu TIM", + "DVD": "DDD", + "Atendimento Sem Controle": "Atendimento TIM Controle", + "MS": "SMS", + "CCB MS": "SMS", + "fiderizado": "fidelizado", + "A P N U S": "A PLUS", + "marga": "carga", + "linha BB": "linha DDD", + "Outubro Premi um": "YouTube Premium", + "Steam Control Smart": "Tim Controle Smart", + "Outubro prêmio": "YouTube Premium", + "e-amil": "e-mail", + "da TIN": "da TIM", + "diesel": "Deezer", + "funeros": "números", + "tendente": "atendente", + "Aitim": "a TIM", + "desastendimento": "atendimento", + "Comprimando": "Confirmando", + "mentinho": "minutinho", + "ilumitada": "ilimitada", + "atendimenta": "atendimento", + "corringir": "corrigir", + "sehora": "senhora", + "planus": "planos", + "cancionamento": "cancelamento", + "atendeimento": "atendimento", + "fielidade": "fidelidade", + "natinho": "na TIM", + "bloqueo": "bloqueio", + "bloqueiro": "bloqueio", + "opição": "opção", + "pagimento": "pagamento", + "reajausto": "reajuste", + "severeir": "fevereiro", + "arrecimento": "vencimento", + "cansele": "cancele", + "cancelimento": "cancelamento", + "cabando": "acabando", + "futbol": "futebol", + "concuído": "concluído", + "imidiato": "imediato", + "aparelo": "aparelho", + "promoción": "promoção", + "outobro": "outubro", + "atrazadas": "atrasadas", + "transferendo": "transferindo", + "pajo": "pago", + "descortos": "descontos", + "sociales": "sociais", + "inoutubra": "outubro", + "deslugar": "desligar", + "cartón": "cartão", + "SNS": "SMS", + "prépaca": "pré pago", + "pessoalamento": "pessoalmente", + "canceloamento": "cancelamento", + "ilumitado": "ilimitado", + "gentilesza": "gentileza", + "promacional": "promocional", + "pás-plago": "pós pago", + "arroma": "arroba", + "contacto": "contato", + "cancola": "cancela", + "benefitios": "benefícios", + "cogrado": "cobrado", + "pasado": "passado", + "fator de comunicação": "falta de comunicação", + "Dizer": "Deezer", + "Dizzer": "Deezer", + "atragado": "atrasado", + "sem CPIN": "sem ser TIM", + "mas como alinhar": "mas como a linha", + "Ei-Fi": "Wi-Fi", + "DND": "DDD", + "CMPJ": "CNPJ", + "fatora": "fatura", + "filmando alguns dados": "confirmando alguns dados", + "Insta gram": "Instagram", + "Aceboca": "Facebook", + "chifre": "chip", + "chave PI": "chave PIX", + "plan": "plano", + "rec carga": "recarga", + "portibilidade": "portabilidade", + "canselamento": "cancelamento", + "cogrou": "cobrou", + "febreiro": "fevereiro", + "pagamenta": "pagamento", + "alinhão": "a linha", + "vo": "vou", + "Aliguação": "A ligação", + "cartõe": "cartões", + "Setembra": "Setembro", + "clientela": "cliente ela", + "primeirinhos": "primeiros dígitos", + "via Pi": "via PIX", + "retaca por aplicativo": "recarga por aplicativo", + "concentração de fatura": "contestação de fatura", + "canal de crédito": "cartão de crédito", + "saldo de leve dor": "saldo devedor", + "a receita da fatura": "a respeito da fatura", + "prescrita de satisfação": "pesquisa de satisfação", + "traca de plano": "troca de plano", + "Conquem, falo": "com quem falo", + "mesmo que vem": "mês que vem", + "se não me enchamos": "se não me engano", + "As terístico": "Asterisco", + "Patinho": "TIM", + "senadora": "senhora", + "promoçal": "promoção", + "jeneiro": "janeiro", + "promoça": "promoção", + "Team black": "TIM Black", + "Team black Cell light": "TIM Black C Light", + "desnoquear": "desbloquear", + "Relígio": "Religa", + "aguardada": "aguardando", + "CIM": "TIM", + "iClap": "iCloud", + "aliar": "avaliar", + "consultra": "consulta", + "do Controlo": "do Controle", + "Batim": "TIM", + "a afilhação caiu": "a ligação caiu", + "a guardinha": "aguarde", + "meu titinho": "Atendimento TIM", + "TeamControlia": "TIM Controle", + "TeamControlia Plus": "TIM Controle Plus", + "Ouvidaria a Atendimento": "Ouvidoria atendimento", + "deputado do": "debitado do", + "Central, te encontrou": "Central TIM Controle", + "botinho da ocorrência": "boletim de ocorrência", + "Bondinha": "bom dia", + "asteristico": "asterisco", + "cum": "com", + "podi transferi": "pode transferir", + "num ti falei": "não te falei", + "não tendi": "não entendi", + "tendi": "entendi", + "UTI": "TIM", + "ubi": "Uber", + "ocilando": "oscilando", + "Team Controle": "TIM Controle", + "Team Controle Smart": "TIM Controle Smart", + "Team Controle A Plus": "TIM Controle A Plus", + "código de barra": "código de barras", + "Team Black C Light": "TIM Black C Light", + "Dá um fatinho. Amém, Tim": "Atendimento TIM", + "do meu PIN": "do meu TIM", + "LÁ Plus” / “La Plus": "Tim Plus", + "código de base": "código de barras", + "código de baixo": "código de barras", + "discólito": "desconto", + "com que a falsa ente lesa": "com quem falo", + "Prefeitura o pagamento": "prefere o pagamento", + "se inscreva no canal": "", + "Seraca": "Serasa", + "Dia PIX": "via PIX", + "fiquei na PIN mesmo": "fiquei na TIM mesmo", + "Ouvidoria é tímido": "Ouvidoria TIM", + "NG": "4G", + "minha vida": "minha linha", + "Ouvidoria, sim": "Ouvidoria TIM", + "ela goa": "Alagoas", + "partura": "fatura", + "Team Controller": "TIM Controle", + "titilaridade": "titularidade", + "Boas tarde": "Boa tarde", + "não sei o que eu falo": "com quem eu falo", + "operora": "operadora", + "Lacoa": "Alagoas", + "ligagem": "ligação", + "canceledado": "cancelado", + "fraud": "fraude", + "Receta": "Receita", + "anteriormento": "anteriormente", + "proprocional": "proporcional", + "procedamento": "procedimento", + "praço": "prazo", + "11gigas": "11 gigas", + "PIM": "TIM", + "ligração": "ligação", + "agorda": "aguarda", + "apareлho": "aparelho", + "senora": "senhora", + "Suha": "sua", + "confirmoção": "confirmação", + "concruir": "concluir", + "fidalidade": "fidelidade", + "correcção": "correção", + "inactivo": "inativo", + "cicl": "ciclo", + "muta": "multa", + "factura": "fatura", + "facturada": "fatura", + "facture": "fatura", + "Novembra": "Novembro", + "Pront": "Pronto", + "faturra": "fatura", + "imprinou": "imprimiu", + "factur": "fatura", + "cancelamenta": "cancelamento", + "opción": "opção", + "cancelamiento": "cancelamento", + "Asterisk": "asterisco", + "ligução": "ligação", + "mudrou": "mudou", + "recalga": "recarga", + "díos": "dias", + "presentidas": "apresentadas", + "bairo": "bairro", + "Team": "TIM", + "pôs-pargo": "pós-pago", + "pó-spago": "pós-pago", + "planhos": "planos", + "setор": "setup", + "fidedilização": "fidelização", + "clientele": "cliente ele", + "disconto": "desconto", + "fidalizado": "fidelizado", + "direcimar": "direcionar", + "cancelamente": "cancelamento", + "beneficios": "benefícios", + "setebro": "setembro", + "iliminada": "eliminada", + "Outubra": "Outubro", + "iliminadas": "eliminadas", + "aletrônico": "eletrônico", + "ligguei": "liguei", + "conclue": "conclui", + "disculpe": "desculpe", + "entrezi": "entendi", + "mensage": "mensagem", + "promorcional": "promocional", + "reajauste": "reajuste", + "TIN": "TIM", + "planas": "planos", + "os panos aumentaram": "os planos aumentaram", + "ocrido": "ocorrido", + "fibrilidade": "fidelidade", + "fidedilidade": "fidelidade", + "sidelidad": "fidelidade", + "sidelizado": "fidelidade", + "fecereiro": "fevereiro", + "vevereira": "fevereiro", + "feveiro": "fevereiro", + "atendimentar": "atendimento", + "pagei": "paguei", + "mismo": "mesmo", + "nômero": "número", + "consulamento": "cancelamento", + "fatural": "fatura", + "salvatura": "sua fatura", + "salvaatura": "sua fatura", + "gigabas": "giga de", + "avalhar": "avaliar", + "pergânticas": "perguntinhas", + "Ateimento": "Atendimento", + "pré-pac": "pré-pago", + "atentimento": "atendimento", + "nuevo": "novo", + "menuna": "menina", + "ontes": "antes", + "Técnoco": "Técnico", + "ontam": "ontem", + "retragir": "retroagir", + "antivalor": "ativação", + "protocoldo": "protocolo", + "Pospag": "Pós-pago", + "Núnmeros": "números", + "aplicativ": "aplicativo", + "perifico": "verifico", + "Ethernet": "internet", + "consigue": "consegue", + "pagado": "pagando", + "namo": "nome", + "gigais": "gigas", + "mudançada": "mudança da", + "distanciaamento": "distanciamento", + "cortada": "contando", + "que eu li aí": "que eu ligo aí", + "bônos": "bônus", + "bónus": "bônus", + "descontro": "desconto", + "septembra": "setembro", + "atingimento": "atendimento", + "pelo favor": "por favor", + "contura": "fatura", + "pacoto": "pacote", + "recargos": "recargas", + "formamento": "atendimento", + "daram": "deram", + "tentiva": "tentativa", + "perifique": "verifiquei", + "faturinha": "fatura", + "cancelar o ciclano": "cancelar esse plano", + "confirmas": "confirmar", + "ligações delimitadas": "ligações ilimitadas", + "Rp": "R$", + "emberto": "aberto", + "receendo": "recebendo", + "fideilização": "fidelização", + "TeamPretop": "TIM Pré Top", + "Prano novo": "Plano novo", + "atendimiento": "atendimento", + "Hoje não": "Pois não", + "meu plano é posto": "meu plano é pós", + "pagote": "pacote", + "oferte": "oferta", + "portbilidade": "portabilidade", + "suportabilidad": "portabilidade", + "chique de atraso": "chip da claro", + "fiz a possibilidade": "fiz a portabilidade", + "gente leza": "gentileza", + "se tem alguma tendência": "se tem alguma pendência", + "a sapatura": "essa fatura", + "Black Day Light": "Black D Light", + "carturas": "faturas", + "Central Outim": "Central TIM", + "código PISC": "código PIX", + "da Tinha": "da TIM", + "home nacional": "roaming nacional", + "Homem internacional": "roaming internacional", + "loja da Tinha": "loja da TIM", + "na TIN": "na TIM", + "número TIN": "número TIM", + "ouvedoria": "ouvidoria", + "para a Tini": "para a TIM", + "Pixi": "PIX", + "PPF": "CPF", + "Presship": "Pré Chip", + "Rendimento Team": "Atendimento Team", + "Simpress Top": "TIM Pre Top", + "Simpresship": "TIM Pre Chip", + "site da Tia": "site da TIM", + "Team Beta": "Tim Beta", + "Team Black": "TIM Black", + "Team Black Alight": "TIM Black A Light", + "Team Black Belight": "Tim Black B Light", + "Team Black Cellite": "Team Black C Light", + "Team Black Delight": "Team Black D Light", + "Team Black Família": "TIM Black Família", + "Team Black Sea Ultra": "Team Black C Ultra", + "Team Black ser ultra": "Team Black C Ultra", + "team família": "TIM Família", + "Team Finanças Mensal": "TIM Finanças Mensal", + "Team pra Team": "TIM pra TIM", + "Team Pre-Chip Plus": "TIM Pre-Chip Plus", + "Team Pre-Chip Plus 1.0": "TIM Pre-Chip Plus 1.0", + "Team Presship Plus": "TIM Pre-Chip Plus", + "Team Pre-Top": "TIM Pre-Top", + "Team Turismo": "TIM Turismo", + "TeamFan": "TIM FUN", + "Tim BlackBelight": "Tim Black B Light", + "Tim BlackC Ultra": "Tim Black C Ultra", + "vencionamento": "vencimento", + "Eu confio": "Eu confirmo", + "Eu confumo": "Eu confirmo", + "Atenso": "Aceito", + "Noun": "Não", + "Nouns": "Não", + "A, X e banca": "AYA, EXA e BANCAH", + "para mal": "Paramount", + "para malte": "Paramount", + "paramalte": "Paramount", + "para monte": "Paramount", + "paramonte": "Paramount", + "paramont": "Paramount", + "para mount": "Paramount", + "para maunt": "Paramount", + "peramount": "Paramount", + "paramounti": "Paramount", + "AIA, Hexa e Banca": "AYA, EXA e BANCAH", + "Hexacloud": "Exa Cloud", + "ExaCloud": "Exa Cloud", + "Tanochi mensal": "Tanoshi Mensal", + "Tanoche mensal": "Tanoshi Mensal", + "Tanois Dimensão": "Tanoshi Mensal", + "TIM Passion": "TIM Fashion", + "TIM Crowd Gaming": "TIM Cloud Gaming", + "TeamCloud Game": "TIM Cloud Gaming", + "PIN Cloud Game": "TIM Cloud Gaming", + "Crowd Gaming": "Cloud Gaming", + "Game Loft": "Gameloft", + "Gabeloft": "Gameloft", + "Instarena": "Insta Arena", + "Instar ainda": "InstArena", + "o book": "Ubook", + "o Bob": "Ubook", + "e book": "Ubook", + "Black Night": "Blacknut", + "Black Nut": "Blacknut", + "Black Nuts": "Blacknut", + "Babel": "Babbel", + "Sem Clube": "Zenklub", + "Zen Clube": "Zenklub", + "Cultive": "Kultivi", + "A WG": "AWG", + "Banca Prime ou Mais": "Bancah Premium", + "Pair Months Mais": "Paramount", + "Permute Mais": "Paramount", + "E Burn": "WeBurn", + "E Born": "WeBurn", + "Ibor": "WeBurn", + "Heading": "Reading", + "Reding": "Reading", + "Buy Lock": "By Looke", + "Luck": "Looke", + "Lut": "Looke", + "Esquilo": "Skeelo", + "Dizze Premium": "Deezer Premium", + "DJ Premium": "Deezer Premium", + "Leo's Gate": "Lionsgate", + "Mônica verso": "Monicaverso", + "Curta 1": "CurtaOn", + "Chefs Club": "ChefsClub", + "Fuso e Forge": "Fuze Forge", + "Mini Boost": "Minibooks", + "Aia e Ensinar": "Aya Ensinah", + "Aia em cima": "Aya Ensinah", + "Zoar em cima": "Aya Ensinah", + "Aion em cima": "Aya Ensinah", + "Aí a box": "Aya Books", + "Que cai a box": "Aya Books", + "Ayabucks": "Aya Books", + "Tem mais topem": "TIM Thopen", + "PIN mas topem": "TIM Thopen" + }, + "max_window": 8, + "exact_first": True, + "full_replacement": True, + "separator_equivalence": True, + "include_metadata": True + } + }, + { + "strategy": "two_stage_hardneg", + "config": { + "artifact_dir": "/usr/src/app/artifacts/two_stage_hardneg", + "device": "auto", + "local_files_only": True, + "transcription_language": "portuguese", + "fail_open": True, + "duration_threshold_ms": 4000 + } + }, + { + "strategy": "clean_default", + "config": { + "regex": "(?i)(?:^\\s*[\\.,!?;:()\\[\\]{}—…\"'`-]+\\s*$|\\b(?:e\\s*a[ií]|ent[aã]o\\s*t[aá]|tchau(?:\\s*[\\.,!?;:()\\[\\]{}—…\"'`-]*\\s*galera)?|valeu|obrigado|se\\s*inscreva\\s*no\\s*canal)\\b[\\.,!?;:()\\[\\]{}—…\"'`-]*)" + } + } + ] + } diff --git a/src/app/providers/stt_fake.py b/src/app/providers/stt_fake.py new file mode 100644 index 0000000..4083a3d --- /dev/null +++ b/src/app/providers/stt_fake.py @@ -0,0 +1,179 @@ +from __future__ import annotations + +import asyncio +import os +import time +import uuid +from typing import Any, Callable, List + +from livekit.agents.stt.stt import ( + STT, + STTCapabilities, + SpeechData, + SpeechEvent, + SpeechEventType, +) +from livekit.agents.types import APIConnectOptions, DEFAULT_API_CONNECT_OPTIONS, NOT_GIVEN, NotGivenOr +from livekit.agents.utils.audio import AudioBuffer, calculate_audio_duration + +from app.utils.logging import ( + EVENT_RECEBIMENTO_MSG, + event_latency_ms, + log_structured_event, + setup_minimal_logging, +) +from app.utils.turn_ids import ( + next_turn_message_id, + register_started_turn_message_id, + register_transcribed_turn, +) + +logger = setup_minimal_logging() + + +def _parse_transcripts(raw_value: str) -> List[str]: + parts: List[str] = [] + for chunk in (raw_value or "").replace("\n", "|").split("|"): + text = chunk.strip() + if text: + parts.append(text) + return parts + + +class FakeSTT(STT): + def __init__( + self, + *, + language: str = "pt-BR", + timeline: Any | None = None, + structured_log_context: Any | None = None, + structured_logger: Any | None = None, + structured_interruption_flag: Callable[[], bool] | None = None, + ) -> None: + super().__init__(capabilities=STTCapabilities(streaming=False, interim_results=False)) + self._language = (language or "pt-BR").strip() or "pt-BR" + self._timeline = timeline + self._structured_log_context = structured_log_context + self._structured_logger = structured_logger + self._structured_interruption_flag = structured_interruption_flag + self._transcripts = _parse_transcripts( + os.getenv( + "FAKE_STT_TRANSCRIPTS", + "alô|quero seguir com a simulacao|sim|obrigado, tchau", + ) + ) + if not self._transcripts: + self._transcripts = ["alô", "sim", "obrigado, tchau"] + self._mode = (os.getenv("FAKE_STT_MODE", "repeat_last") or "repeat_last").strip().lower() + self._min_audio_ms = max(0, int(os.getenv("FAKE_STT_MIN_AUDIO_MS", "150") or "150")) + self._cursor = 0 + self._lock = asyncio.Lock() + + def _structured_interruption_value(self) -> int: + if self._structured_interruption_flag is None: + return 0 + try: + return 1 if bool(self._structured_interruption_flag()) else 0 + except Exception: + return 0 + + @property + def model(self) -> str: + return "fake-script" + + @property + def provider(self) -> str: + return "fake" + + async def _recognize_impl( + self, + buffer: AudioBuffer, + *, + language: NotGivenOr[str] = NOT_GIVEN, + conn_options: APIConnectOptions = DEFAULT_API_CONNECT_OPTIONS, + ) -> SpeechEvent: + del conn_options + + started_ns = time.time_ns() + request_id = uuid.uuid4().hex + audio_duration_ms = round(calculate_audio_duration(buffer) * 1000) + resolved_language = language if isinstance(language, str) and language.strip() else self._language + + if audio_duration_ms < self._min_audio_ms: + if self._timeline is not None: + self._timeline.emit( + "stt_fake_skipped", + request_id=request_id, + duration_ms=audio_duration_ms, + reason="too_short", + ) + return SpeechEvent( + type=SpeechEventType.FINAL_TRANSCRIPT, + request_id=request_id, + alternatives=[SpeechData(language=resolved_language, text="", confidence=0.0)], + ) + + async with self._lock: + current_index = self._cursor + if self._mode == "cycle": + transcript = self._transcripts[current_index % len(self._transcripts)] + self._cursor += 1 + else: + bounded_index = min(current_index, len(self._transcripts) - 1) + transcript = self._transcripts[bounded_index] + if self._cursor < len(self._transcripts) - 1: + self._cursor += 1 + + if transcript.lower() in {"__silence__", "__empty__"}: + transcript = "" + + if self._timeline is not None: + self._timeline.emit( + "stt_fake_completed", + request_id=request_id, + duration_ms=audio_duration_ms, + transcript=transcript, + transcript_index=current_index, + ) + if transcript: + finished_ns = time.time_ns() + logger.info( + "[stt][fake_transcript] req_id=%s duration_ms=%s audio_duration_ms=%s text=%r", + request_id, + event_latency_ms(started_ns, finished_ns), + audio_duration_ms, + transcript, + ) + message_id = "" + if transcript: + message_id = next_turn_message_id(self._structured_log_context) + register_started_turn_message_id(self._structured_log_context, message_id) + register_transcribed_turn( + self._structured_log_context, + message_id=message_id, + transcription=transcript, + text=transcript, + ) + + if transcript and self._structured_logger is not None: + total_ms = event_latency_ms(started_ns, finished_ns) + log_structured_event( + self._structured_logger, + self._structured_log_context, + tipo_evento=EVENT_RECEBIMENTO_MSG, + message_id=message_id, + inicio_ns=started_ns, + fim_ns=finished_ns, + latencia_total_ms=total_ms, + latencia_tffb_ms=total_ms, + interrupcao=self._structured_interruption_value(), + ) + + return SpeechEvent( + type=SpeechEventType.FINAL_TRANSCRIPT, + request_id=request_id, + alternatives=[SpeechData(language=resolved_language, text=transcript, confidence=1.0)], + ) + + async def aclose(self) -> None: + return None diff --git a/src/app/providers/stt_internal_livekit.py b/src/app/providers/stt_internal_livekit.py new file mode 100644 index 0000000..ae37633 --- /dev/null +++ b/src/app/providers/stt_internal_livekit.py @@ -0,0 +1,1180 @@ +from __future__ import annotations + +import re +import unicodedata +import asyncio +import copy +import threading +import os +import io +import json +import uuid +import wave +import time +import logging +from dataclasses import dataclass +from typing import Any, Callable, Dict, Optional, Tuple + +import httpx +from livekit import rtc +from livekit.agents.stt.stt import ( + STT, + STTCapabilities, + SpeechEvent, + SpeechEventType, + SpeechData, + APIConnectOptions, +) + +from app.utils.pcm import pcm_duration_ms, dbfs_pcm16le +from app.utils.dump import _dump_audio_for_debug +from app.utils.stt_audio_upload import enqueue_stt_vad_audio_upload +from app.providers.stt_config import override_config +from app.providers.stt_vosk import VoskSTT +from app.providers.stt_word_dictionary import ( + SHORT_WORD_ALLOWLIST, + SHORT_WORD_ALLOWLIST_MIN_PROB, +) +from app.common.timed import timed +from app.utils.logging import ( + EVENT_RECEBIMENTO_MSG, + event_latency_ms, + log_flow_event, + log_structured_event, + setup_minimal_logging, +) +from app.utils.turn_ids import ( + clear_started_turn_message_id, + next_turn_message_id, + register_started_turn_message_id, + register_transcribed_turn, +) + +logger = setup_minimal_logging() + +_HTTP_500_MAX_RETRIES = 1 + +# --- single-word rewrite (apenas para raw_json) --- +_SINGLE_WORD_REWRITE_MAP = { + "bom": "Não", + "bem": "Sim", + "falou": "Alô", +} + +_FINAL_PUNCT_RE = re.compile(r"[\s]*[.!?…,:;]+[\s]*$") + +def _env_non_negative_int(name: str, default: int) -> int: + try: + return max(0, int(os.getenv(name, str(default)))) + except ValueError: + return max(0, default) + + +def _json_for_log(payload: Any) -> str: + try: + return json.dumps(payload, ensure_ascii=False, default=str) + except Exception: + return str(payload) + + +def _prepend_pcm16le_silence( + pcm: bytes, + *, + sample_rate: int, + channels: int, + padding_ms: int, +) -> tuple[bytes, int]: + if not pcm or padding_ms <= 0: + return pcm, 0 + sample_frames = round(sample_rate * (padding_ms / 1000.0)) + silence = b"\x00" * max(0, sample_frames) * max(1, channels) * 2 + if not silence: + return pcm, 0 + actual_padding_ms = round(pcm_duration_ms(silence, sample_rate, channels)) + return silence + pcm, actual_padding_ms + + +def strip_final_punctuation(text: str) -> str: + return _FINAL_PUNCT_RE.sub("", (text or "").strip()) + +def _normalize_token(s: str) -> str: + s = strip_final_punctuation(s) # remove pontuação final + s = (s or "").strip().lower() + s = unicodedata.normalize("NFD", s) + s = "".join(ch for ch in s if unicodedata.category(ch) != "Mn") # remove acentos + return s + + +_NORMALIZED_SHORT_WORD_ALLOWLIST = frozenset( + _normalize_token(word) + for word in SHORT_WORD_ALLOWLIST + if _normalize_token(word) +) + + +def _is_short_word_allowlisted(text: str) -> bool: + return _normalize_token(text) in _NORMALIZED_SHORT_WORD_ALLOWLIST + + +def _short_word_allowlist_accepts(text: str, _probability: float) -> bool: + return _is_short_word_allowlisted(text) + + +def apply_single_word_rewrite_for_raw_json(norm_payload: Dict[str, Any]) -> Dict[str, Any]: + """ + Se o payload (já normalizado em {"data":{"text","words"}}) tiver 1 palavra, + troca data.text conforme _SINGLE_WORD_REWRITE_MAP. + """ + try: + data = norm_payload.get("data") if isinstance(norm_payload.get("data"), dict) else {} + text = (data.get("text") or "").strip() + words = data.get("words") if isinstance(data.get("words"), list) else [] + + # "retorno for 1": prioriza words.length, senão cai pra contagem de tokens do text + if words: + is_single = (len(words) == 1) + else: + is_single = (len([t for t in _normalize_token(text).split() if t]) == 1) + + if not is_single or not text: + return norm_payload + + key = _normalize_token(text) + replacement = _SINGLE_WORD_REWRITE_MAP.get(key) + if replacement: + data["text"] = replacement + norm_payload["data"] = data + + # Se quiser manter consistência com words[0]["word"], descomente: + # if isinstance(words, list) and len(words) == 1 and isinstance(words[0], dict): + # words[0]["word"] = f" {replacement}" + # data["words"] = words + + return norm_payload + except Exception: + return norm_payload + + +def extract_text_from_transcript(transcript: str) -> str: + """ + Aceita: + - texto puro: "oi tudo bem" + - JSON string novo: {"data":{"text":"...","words":[...]}} + - JSON string legado: {"text":"..."} ou {"data":{"text":"..."}} + + Retorna SEMPRE texto limpo (ou "" se não achar texto). + """ + s = (transcript or "").strip() + if not s: + return "" + + if s[:1] in "{[": + try: + obj = json.loads(s) + + if isinstance(obj, dict): + data = obj.get("data") + if isinstance(data, dict): + t = data.get("text") + if isinstance(t, str) and t.strip(): + return t.strip() + + t = obj.get("text") + if isinstance(t, str) and t.strip(): + return t.strip() + + return "" + + if isinstance(obj, list): + for it in obj: + if isinstance(it, dict): + data = it.get("data") + if isinstance(data, dict): + t = data.get("text") + if isinstance(t, str) and t.strip(): + return t.strip() + t = it.get("text") + if isinstance(t, str) and t.strip(): + return t.strip() + return "" + except Exception: + pass + + return s + + +def normalize_to_data_words_format(payload: Dict[str, Any]) -> Dict[str, Any]: + """ + Normaliza diferentes formatos para: + {"data": {"text": "...", "words": [{"word","start","end","probability"}]}} + + Suporta: + - {"data": {"text": "...", "words":[{"probability":...}]}} + - {"text": "...", "words":[...]} (move pra dentro de data) + - {"text": "...", "result":[{"word","start","end","conf"}]} (Vosk padrão) + - {"text": "...", "Probabilites"/"Probabilities":[{"words":[...]}]} (legado) + """ + if not isinstance(payload, dict): + return {"data": {"text": "", "words": []}} + + # já é novo + if isinstance(payload.get("data"), dict): + data = payload["data"] + text = (data.get("text") or "").strip() + words = data.get("words") if isinstance(data.get("words"), list) else [] + return {"data": {"text": text, "words": words}} + + root = payload + text = (root.get("text") or "").strip() + + # words direto + if isinstance(root.get("words"), list): + return {"data": {"text": text, "words": root["words"]}} + + # vosk "result" + if isinstance(root.get("result"), list): + out_words = [] + for w in root["result"]: + if not isinstance(w, dict): + continue + out_words.append({ + "word": w.get("word", ""), + "start": float(w.get("start", 0.0) or 0.0), + "end": float(w.get("end", 0.0) or 0.0), + "probability": float(w.get("conf", w.get("probability", 0.0)) or 0.0), + }) + return {"data": {"text": text, "words": out_words}} + + # legado Probabilites/Probabilities + probs = ( + root.get("Probabilites") + or root.get("Probabilities") + or root.get("probabilites") + or root.get("probabilities") + or [] + ) + if isinstance(probs, list) and probs: + first = probs[0] if isinstance(probs[0], dict) else {} + words = first.get("words") if isinstance(first.get("words"), list) else [] + seg_text = (first.get("text") or text or "").strip() + return {"data": {"text": seg_text, "words": words}} + + return {"data": {"text": text, "words": []}} + + +def stt_text_with_single_word_threshold( + payload: Dict[str, Any], + min_prob_single_word: float = 0.10, +) -> Optional[str]: + """ + Regra: + - allowlist passa independente da probabilidade + - 1 palavra: só aceita se prob > min_prob_single_word + - >1 palavra: aceita + - sem words: aceita text + """ + try: + norm = normalize_to_data_words_format(payload) + data = norm.get("data") if isinstance(norm.get("data"), dict) else {} + text = (data.get("text") or "").strip() + if not text: + return None + + words = data.get("words") if isinstance(data.get("words"), list) else [] + if not words: + return text + + if len(words) == 1: + prob = words[0].get("probability", 0.0) + try: + prob = float(prob) + except Exception: + prob = 0.0 + if prob > min_prob_single_word: + return text + if _short_word_allowlist_accepts(text, prob): + return text + return None + + return text + except Exception: + return None + + +@dataclass +class InternalSTTConfig: + url: str + api_key: str + language: str = "pt" + sample_rate: int = 16000 + channels: int = 1 + timeout_s: float = 30.0 + min_prob_single_word: float = 0.10 + initial_prompt: str = "" + config_override: str = "" + output_mode: str = "threshold_text" # "threshold_text" | "raw_json" | "api_text" + + +def _normalize_override_config_json(raw_config: str) -> str: + text = (raw_config or "").strip() + if not text: + return "" + + payload = json.loads(text) + if not isinstance(payload, dict): + raise ValueError("STT config override must be a JSON object") + return json.dumps(payload, ensure_ascii=False) + + +class InternalHTTPSTT(STT): + """ + STT batch (não-streaming). + O LiveKit Agents vai usar VAD/turn_detection para segmentar e chamar recognize() por trecho. + """ + + def __init__( + self, + cfg: InternalSTTConfig, + *, + client: httpx.AsyncClient, + vosk: Optional[VoskSTT] = None, + vosk_lock: Optional[threading.Lock] = None, + timeline: Any | None = None, + structured_log_context: Any | None = None, + structured_logger: Any | None = None, + structured_interruption_flag: Callable[[], bool] | None = None, + empty_transcript_handler: Callable[[], None] | None = None, + metrics_handler: Callable[[dict[str, Any]], None] | None = None, + ): + super().__init__(capabilities=STTCapabilities(streaming=False, interim_results=False)) + self._cfg = cfg + self._client = client + self._vosk = vosk + self._vosk_lock = vosk_lock + self._timeline = timeline + self._structured_log_context = structured_log_context + self._structured_logger = structured_logger + self._structured_interruption_flag = structured_interruption_flag + self._empty_transcript_handler = empty_transcript_handler + self._metrics_handler = metrics_handler + self._resampler: Optional[rtc.AudioResampler] = None + + def set_empty_transcript_handler(self, handler: Callable[[], None] | None) -> None: + self._empty_transcript_handler = handler + + def set_metrics_handler( + self, handler: Callable[[dict[str, Any]], None] | None + ) -> None: + self._metrics_handler = handler + + def _notify_metrics(self, **fields: Any) -> None: + handler = self._metrics_handler + if handler is None: + return + try: + handler(fields) + except Exception: + logger.exception("[stt][metrics_handler_error]") + + def _notify_empty_transcript(self, *, request_id: str, source: str) -> None: + handler = self._empty_transcript_handler + if handler is None: + return + try: + handler() + except Exception: + logger.exception( + "[stt][empty_transcript_handler_error] req_id=%s source=%s", + request_id, + source, + ) + + def _structured_interruption_value(self) -> int: + if self._structured_interruption_flag is None: + return 0 + try: + return 1 if bool(self._structured_interruption_flag()) else 0 + except Exception: + return 0 + + def _log_recebimento_msg( + self, + *, + request_id: str, + started_ns: int, + message_id: str = "", + finished_ns: int | None = None, + erro_msg: str | None = None, + erro_detalhe: str | None = None, + http_cod_status: int | None = None, + http_cod_desc: str | None = None, + ) -> str: + resolved_message_id = str(message_id or "").strip() or next_turn_message_id( + self._structured_log_context + ) + register_started_turn_message_id(self._structured_log_context, resolved_message_id) + if self._structured_logger is None: + return resolved_message_id + + ended_ns = int(finished_ns if finished_ns is not None else time.time_ns()) + total_ms = event_latency_ms(started_ns, ended_ns) + log_structured_event( + self._structured_logger, + self._structured_log_context, + tipo_evento=EVENT_RECEBIMENTO_MSG, + message_id=resolved_message_id, + inicio_ns=started_ns, + fim_ns=ended_ns, + latencia_total_ms=total_ms, + latencia_tffb_ms=total_ms, + interrupcao=self._structured_interruption_value(), + erro_msg=erro_msg, + erro_detalhe=erro_detalhe, + http_cod_status=http_cod_status, + http_cod_desc=http_cod_desc, + ) + return resolved_message_id + + def _log_detected_text( + self, + *, + request_id: str, + output_text: str, + engine: str, + mode: str | None = None, + duration_ms: float | int | None = None, + audio_duration_ms: float | int | None = None, + level_dbfs: float | None = None, + ) -> None: + detected_text = extract_text_from_transcript(output_text) + log_flow_event( + self._structured_logger or logger, + "stt_done", + request_id=request_id, + engine=engine, + mode=mode or self._stt_output_mode(), + returned=True, + duration_ms=round(float(duration_ms)) if duration_ms is not None else None, + audio_ms=round(float(audio_duration_ms)) if audio_duration_ms is not None else None, + dbfs=round(float(level_dbfs), 2) if level_dbfs is not None else None, + text_len=len(detected_text), + text=detected_text, + ) + + def _register_transcribed_turn(self, *, message_id: str, transcription: str) -> None: + detected_text = extract_text_from_transcript(transcription) + if not detected_text: + return + + register_transcribed_turn( + self._structured_log_context, + message_id=message_id, + transcription=transcription, + text=detected_text, + ) + + def _stt_output_mode(self) -> str: + return (os.getenv("STT_OUTPUT_MODE", self._cfg.output_mode) or "threshold_text").strip().lower() + + def _format_stt_output(self, payload: Dict[str, Any], *, request_id: str = "") -> str: + """ + Converte o payload para uma string, dependendo do modo: + + - raw_json: devolve JSON string no formato {"data":{...}} + - api_text: devolve apenas data.text + - threshold_text: aplica regra de 1 palavra com prob + """ + mode = self._stt_output_mode() + + norm = normalize_to_data_words_format(payload) + data = norm.get("data") if isinstance(norm.get("data"), dict) else {} + api_text = (data.get("text") or "").strip() + words = data.get("words") if isinstance(data.get("words"), list) else [] + single_word_prob = None + if len(words) == 1 and isinstance(words[0], dict): + try: + single_word_prob = float(words[0].get("probability", 0.0) or 0.0) + except Exception: + single_word_prob = 0.0 + + if mode == "raw_json": + logger.info("[stt][raw_json_before_rewrite] %s", norm) + norm = apply_single_word_rewrite_for_raw_json(norm) + norm["STT"] = "Sofya" + logger.info("[stt][raw_json_after_rewrite] %s", norm) + return json.dumps(norm, ensure_ascii=False) + + + if mode == "api_text": + return (data.get("text") or "").strip() + + out_text = stt_text_with_single_word_threshold(norm, self._cfg.min_prob_single_word) or "" + filter_reason = "accepted" + if ( + out_text + and len(words) == 1 + and single_word_prob is not None + and single_word_prob <= self._cfg.min_prob_single_word + and _short_word_allowlist_accepts(api_text, single_word_prob) + ): + filter_reason = "single_word_allowlist_low_confidence" + if not out_text: + if not api_text: + filter_reason = "api_text_empty" + elif len(words) == 1 and single_word_prob is not None: + filter_reason = "single_word_below_threshold" + else: + filter_reason = "threshold_text_empty" + + log_flow_event( + self._structured_logger or logger, + "stt_payload", + request_id=request_id or None, + mode=mode, + api_text=api_text, + api_text_len=len(api_text), + words_count=len(words), + min_prob_single_word=self._cfg.min_prob_single_word, + single_word_prob=( + f"{single_word_prob:.3f}" if single_word_prob is not None else None + ), + allowlist_min_prob=SHORT_WORD_ALLOWLIST_MIN_PROB, + filter_reason=filter_reason, + ) + return out_text + + def _build_override_config_json(self) -> str: + forced_override = (os.getenv("STT_FORCE_OVERRIDE_CONFIG", "") or "").strip() + if forced_override: + return _normalize_override_config_json(forced_override) + + inline_override = (self._cfg.config_override or "").strip() + if inline_override: + return _normalize_override_config_json(inline_override) + + base = override_config + ovr = copy.deepcopy(base) + + proc = ovr.setdefault("processor", {}).setdefault("config", {}).setdefault("extra_params", {}) + + forced = (os.getenv("STT_FORCE_INITIAL_PROMPT", "").strip() or "").strip() + prompt = forced or (self._cfg.initial_prompt or "").strip() or "Aguarda confimação de números do cpf" + proc["initial_prompt"] = prompt + + return json.dumps(ovr, ensure_ascii=False) + + async def _try_vosk(self, pcm: bytes, req_id: str) -> Tuple[str, Dict[str, Any]]: + """ + Roda Vosk em thread e serializa com lock. + Retorna: (texto, payload_padronizado_em_data) + """ + if self._vosk is None: + return "", {} + if self._timeline is not None: + self._timeline.emit( + "stt_vosk_started", + request_id=req_id, + pcm_bytes=len(pcm), + ) + + lock = self._vosk_lock + + def _run() -> Tuple[str, Dict[str, Any]]: + if lock is not None: + with lock: + raw = self._vosk.transcribe_pcm16le(pcm) + else: + raw = self._vosk.transcribe_pcm16le(pcm) + + # raw já vem padronizado se seu VoskSTT estiver atualizado, + # mas garantimos aqui também: + txt, payload = VoskSTT.result_to_text_and_payload(raw) + payload = normalize_to_data_words_format(payload) + return (txt or "").strip(), payload + + try: + text, payload = await asyncio.to_thread(_run) + if self._timeline is not None: + self._timeline.emit( + "stt_vosk_completed", + request_id=req_id, + text=text, + text_len=len(text), + has_payload=bool(payload), + ) + return text, payload + except Exception: + logger.exception("[stt][vosk_error] req_id=%s", req_id) + if self._timeline is not None: + self._timeline.emit("stt_vosk_failed", request_id=req_id) + return "", {} + + async def _post_internal_http( + self, + pcm: bytes, + wav_bytes: bytes, + req_id: str, + started_ns: int, + language: Optional[str], + message_id: str = "", + audio_duration_ms: float | int | None = None, + original_audio_duration_ms: float | int | None = None, + input_padding_ms: float | int | None = None, + level_dbfs: float | None = None, + ) -> SpeechEvent: + files = {"file": ("audio.wav", wav_bytes, "audio/wav")} + headers = {"x-api-key": self._cfg.api_key, "Connection": "close"} + data = {"override_config": self._build_override_config_json()} + timeout = httpx.Timeout(self._cfg.timeout_s) + + if self._timeline is not None: + self._timeline.emit( + "stt_http_started", + request_id=req_id, + url=self._cfg.url, + ) + log_flow_event( + self._structured_logger or logger, + "stt_dispatch", + request_id=req_id, + engine="http", + audio_ms=round(float(audio_duration_ms)) if audio_duration_ms is not None else None, + dbfs=round(float(level_dbfs), 2) if level_dbfs is not None else None, + ) + attempt = 0 + try: + while True: + t0 = time.perf_counter() + r = await self._client.post( + self._cfg.url, + files=files, + data=data, + headers=headers, + timeout=timeout, + ) + dt_ms = (time.perf_counter() - t0) * 1000.0 + try: + r.raise_for_status() + finished_ns = time.time_ns() + break + except httpx.HTTPStatusError as e: + http_status = getattr(e.response, "status_code", None) + if http_status == 500 and attempt < _HTTP_500_MAX_RETRIES: + attempt += 1 + logger.warning( + "[stt][http_500_retry] req_id=%s attempt=%s/%s took=%.0fms url=%s", + req_id, + attempt, + _HTTP_500_MAX_RETRIES, + dt_ms, + str(getattr(e.request, "url", self._cfg.url)), + ) + log_flow_event( + self._structured_logger or logger, + "stt_http_retry", + request_id=req_id, + engine="http", + status=http_status, + attempt=attempt, + max_retries=_HTTP_500_MAX_RETRIES, + duration_ms=round(dt_ms), + ) + if self._timeline is not None: + self._timeline.emit( + "stt_http_retry", + request_id=req_id, + status=http_status, + attempt=attempt, + max_retries=_HTTP_500_MAX_RETRIES, + duration_ms=round(dt_ms), + ) + continue + raise + except httpx.HTTPStatusError as e: + dt_ms = (time.perf_counter() - t0) * 1000.0 + http_status = getattr(e.response, "status_code", None) + http_desc = str(getattr(e.response, "reason_phrase", "") or "").strip() or None + body_preview = "" + try: + body_preview = (e.response.text or "")[:800] + except Exception: + pass + logger.error( + "[stt][http_error] req_id=%s status=%s took=%.0fms url=%s body(800)=%s", + req_id, + getattr(e.response, "status_code", "NA"), + dt_ms, + str(getattr(e.request, "url", self._cfg.url)), + body_preview, + ) + log_flow_event( + self._structured_logger or logger, + "stt_error", + request_id=req_id, + engine="http", + status=http_status or "NA", + duration_ms=round(dt_ms), + ) + if self._timeline is not None: + self._timeline.emit( + "stt_http_failed", + request_id=req_id, + status=getattr(e.response, "status_code", "NA"), + duration_ms=round(dt_ms), + ) + self._log_recebimento_msg( + request_id=req_id, + started_ns=started_ns, + message_id=message_id, + erro_msg="Falha STT", + erro_detalhe=( + f"http_status={http_status if http_status is not None else 'NA'} " + f"| url={str(getattr(e.request, 'url', self._cfg.url))} " + f"| body={body_preview}" + ), + http_cod_status=http_status, + http_cod_desc=http_desc, + ) + self._notify_metrics( + event="failed", + request_id=req_id, + provider="internal_http", + duration_ms=round(dt_ms), + audio_duration_ms=round(float(audio_duration_ms)) if audio_duration_ms is not None else None, + original_audio_duration_ms=( + round(float(original_audio_duration_ms)) + if original_audio_duration_ms is not None + else None + ), + input_padding_ms=round(float(input_padding_ms)) if input_padding_ms is not None else None, + input_dbfs=round(float(level_dbfs), 2) if level_dbfs is not None else None, + retry_count=attempt, + http_status=http_status, + error="http_error", + ) + if http_status == 500: + log_flow_event( + self._structured_logger or logger, + "stt_error_nonfatal", + request_id=req_id, + engine="http", + status=http_status, + action="return_empty_transcript", + max_retries=_HTTP_500_MAX_RETRIES, + ) + if self._timeline is not None: + self._timeline.emit( + "stt_http_failed_nonfatal", + request_id=req_id, + status=http_status, + max_retries=_HTTP_500_MAX_RETRIES, + ) + return SpeechEvent( + type=SpeechEventType.FINAL_TRANSCRIPT, + request_id=req_id, + alternatives=[ + SpeechData( + language=language or self._cfg.language, + text="", + confidence=0.0, + ) + ], + ) + raise + except Exception: + dt_ms = (time.perf_counter() - t0) * 1000.0 + logger.exception("[stt][error] req_id=%s took=%.0fms", req_id, dt_ms) + log_flow_event( + self._structured_logger or logger, + "stt_error", + request_id=req_id, + engine="http", + status="exception", + duration_ms=round(dt_ms), + ) + if self._timeline is not None: + self._timeline.emit( + "stt_http_failed", + request_id=req_id, + status="exception", + duration_ms=round(dt_ms), + ) + self._log_recebimento_msg( + request_id=req_id, + started_ns=started_ns, + message_id=message_id, + erro_msg="Falha STT", + erro_detalhe=f"exception | took={round(dt_ms)}ms", + ) + self._notify_metrics( + event="failed", + request_id=req_id, + provider="internal_http", + duration_ms=round(dt_ms), + audio_duration_ms=round(float(audio_duration_ms)) if audio_duration_ms is not None else None, + original_audio_duration_ms=( + round(float(original_audio_duration_ms)) + if original_audio_duration_ms is not None + else None + ), + input_padding_ms=round(float(input_padding_ms)) if input_padding_ms is not None else None, + input_dbfs=round(float(level_dbfs), 2) if level_dbfs is not None else None, + retry_count=attempt, + http_status=None, + error="exception", + ) + raise + + payload = r.json() + payload_json = _json_for_log(payload) + logger.info( + "[stt][response_json] req_id=%s status=%s took=%.0fms json=%s", + req_id, + r.status_code, + dt_ms, + payload_json, + ) + out_text = self._format_stt_output(payload, request_id=req_id) + if self._timeline is not None: + self._timeline.emit( + "stt_http_completed", + request_id=req_id, + duration_ms=round(dt_ms), + text=out_text, + text_len=len(out_text), + ) + empty_transcript = not bool(extract_text_from_transcript(out_text)) + self._notify_metrics( + event="completed", + request_id=req_id, + provider="internal_http", + mode=self._stt_output_mode(), + duration_ms=round(dt_ms), + audio_duration_ms=round(float(audio_duration_ms)) if audio_duration_ms is not None else None, + original_audio_duration_ms=( + round(float(original_audio_duration_ms)) + if original_audio_duration_ms is not None + else None + ), + input_padding_ms=round(float(input_padding_ms)) if input_padding_ms is not None else None, + input_dbfs=round(float(level_dbfs), 2) if level_dbfs is not None else None, + retry_count=attempt, + http_status=r.status_code, + empty_transcript=empty_transcript, + text_length=len(extract_text_from_transcript(out_text)), + ) + if empty_transcript: + # Um resultado sem texto nao constitui turno de conversa. Nao publique + # recebimento msg nem deixe o identificador tecnico do upload ocupar o + # estado de message_id ativo, pois uma fala proativa do agente poderia + # herda-lo antes da chegada de user_input_transcribed. + clear_started_turn_message_id( + self._structured_log_context, + message_id=message_id, + ) + else: + message_id = self._log_recebimento_msg( + request_id=req_id, + started_ns=started_ns, + message_id=message_id, + finished_ns=finished_ns, + erro_detalhe=None, + http_cod_status=r.status_code, + http_cod_desc=str(getattr(r, "reason_phrase", "") or "").strip() or None, + ) + self._register_transcribed_turn(message_id=message_id, transcription=out_text) + self._log_detected_text( + request_id=req_id, + output_text=out_text, + engine="http", + duration_ms=dt_ms, + audio_duration_ms=audio_duration_ms, + level_dbfs=level_dbfs, + ) + if empty_transcript: + self._notify_empty_transcript(request_id=req_id, source="http_success") + + # log reduzido + preview = out_text if len(out_text) < 250 else out_text[:250] + "..." + + return SpeechEvent( + type=SpeechEventType.FINAL_TRANSCRIPT, + request_id=req_id, + alternatives=[SpeechData(language=language or self._cfg.language, text=out_text, confidence=0.0)], + ) + + async def aclose(self) -> None: + return + + def _frames_to_pcm16le_16k(self, buffer) -> bytes: + """ + Junta frames e garante PCM16 mono 16k (re-amostra se precisar). + """ + if isinstance(buffer, rtc.AudioFrame): + frames = [buffer] + else: + frames = list(buffer) + + if not frames: + return b"" + + in_sr = frames[0].sample_rate + in_ch = frames[0].num_channels + + target_sr = self._cfg.sample_rate + target_ch = self._cfg.channels + + # resample se necessário + if in_sr != target_sr: + if self._resampler is None: + self._resampler = rtc.AudioResampler(in_sr, target_sr, quality=rtc.AudioResamplerQuality.HIGH) + out_frames = [] + for f in frames: + out_frames.extend(self._resampler.push(f)) + out_frames.extend(self._resampler.flush()) + frames = out_frames + + # canais (se precisar no futuro, faz stereo->mono aqui) + if in_ch != target_ch: + pass + + return b"".join(bytes(f.data) for f in frames) + + @timed("STT") + async def _recognize_impl( + self, + buffer, + *, + language: str | None, + conn_options: APIConnectOptions, + ) -> SpeechEvent: + started_ns = time.time_ns() + pcm = self._frames_to_pcm16le_16k(buffer) + req_id = uuid.uuid4().hex + + min_ms = int(os.getenv("STT_MIN_AUDIO_MS", "160")) + min_dbfs = float(os.getenv("STT_MIN_DBFS", "-50.0")) + + dur_ms = pcm_duration_ms(pcm, self._cfg.sample_rate, self._cfg.channels) + level_dbfs = dbfs_pcm16le(pcm) + if self._timeline is not None: + self._timeline.emit( + "stt_recognize_started", + request_id=req_id, + duration_ms=round(dur_ms), + level_dbfs=round(level_dbfs, 2), + ) + if not pcm or dur_ms < min_ms or level_dbfs < min_dbfs: + if (os.getenv("FLOW_LOG_STT_SKIPS", "0") or "0").strip() == "1": + log_flow_event( + self._structured_logger or logger, + "stt_skip", + request_id=req_id, + reason="too_short_or_silent", + audio_ms=round(dur_ms), + dbfs=round(level_dbfs, 2), + ) + if self._timeline is not None: + self._timeline.emit( + "stt_recognize_skipped", + request_id=req_id, + reason="too_short_or_silent", + duration_ms=round(dur_ms), + level_dbfs=round(level_dbfs, 2), + ) + return SpeechEvent( + type=SpeechEventType.FINAL_TRANSCRIPT, + request_id=req_id, + alternatives=[SpeechData(language=language or self._cfg.language, text="", confidence=0.0)], + ) + log_flow_event( + self._structured_logger or logger, + "stt_start", + request_id=req_id, + audio_ms=round(dur_ms), + dbfs=round(level_dbfs, 2), + ) + + # tenta Vosk para trechos curtos + vosk_max_ms = int(os.getenv("STT_VOSK_MAX_MS", "2000")) + if dur_ms < vosk_max_ms: + vosk_started = time.perf_counter() + txt_vosk, payload_vosk = await self._try_vosk(pcm, req_id=req_id) + vosk_duration_ms = (time.perf_counter() - vosk_started) * 1000.0 + + if payload_vosk: + mode = self._stt_output_mode() + + if mode == "raw_json": + payload_vosk = apply_single_word_rewrite_for_raw_json(payload_vosk) + payload_vosk["STT"] = "Vosk" + out_text = json.dumps(payload_vosk, ensure_ascii=False) + if self._timeline is not None: + self._timeline.emit( + "stt_recognize_completed", + request_id=req_id, + engine="vosk", + mode=mode, + text=out_text, + text_len=len(out_text), + ) + message_id = self._log_recebimento_msg( + request_id=req_id, + started_ns=started_ns, + ) + self._register_transcribed_turn(message_id=message_id, transcription=out_text) + self._log_detected_text( + request_id=req_id, + output_text=out_text, + engine="vosk", + mode=mode, + duration_ms=vosk_duration_ms, + audio_duration_ms=dur_ms, + level_dbfs=level_dbfs, + ) + return SpeechEvent( + type=SpeechEventType.FINAL_TRANSCRIPT, + request_id=req_id, + alternatives=[SpeechData(language=language or self._cfg.language, text=out_text, confidence=0.0)], + ) + + if mode == "api_text": + data = payload_vosk.get("data") if isinstance(payload_vosk.get("data"), dict) else {} + out_text = (data.get("text") or "").strip() + if out_text: + if self._timeline is not None: + self._timeline.emit( + "stt_recognize_completed", + request_id=req_id, + engine="vosk", + mode=mode, + text=out_text, + text_len=len(out_text), + ) + message_id = self._log_recebimento_msg( + request_id=req_id, + started_ns=started_ns, + ) + self._register_transcribed_turn(message_id=message_id, transcription=out_text) + self._log_detected_text( + request_id=req_id, + output_text=out_text, + engine="vosk", + mode=mode, + duration_ms=vosk_duration_ms, + audio_duration_ms=dur_ms, + level_dbfs=level_dbfs, + ) + return SpeechEvent( + type=SpeechEventType.FINAL_TRANSCRIPT, + request_id=req_id, + alternatives=[SpeechData(language=language or self._cfg.language, text=out_text, confidence=0.0)], + ) + + # threshold_text + filtered = stt_text_with_single_word_threshold(payload_vosk, self._cfg.min_prob_single_word) + if filtered: + if self._timeline is not None: + self._timeline.emit( + "stt_recognize_completed", + request_id=req_id, + engine="vosk", + mode="threshold_text", + text=filtered, + text_len=len(filtered), + ) + message_id = self._log_recebimento_msg( + request_id=req_id, + started_ns=started_ns, + ) + self._register_transcribed_turn(message_id=message_id, transcription=filtered) + self._log_detected_text( + request_id=req_id, + output_text=filtered, + engine="vosk", + mode="threshold_text", + duration_ms=vosk_duration_ms, + audio_duration_ms=dur_ms, + level_dbfs=level_dbfs, + ) + return SpeechEvent( + type=SpeechEventType.FINAL_TRANSCRIPT, + request_id=req_id, + alternatives=[SpeechData(language=language or self._cfg.language, text=filtered, confidence=0.0)], + ) + + input_padding_ms = _env_non_negative_int("STT_INPUT_PREFIX_PADDING_MS", 250) + http_pcm, applied_input_padding_ms = _prepend_pcm16le_silence( + pcm, + sample_rate=self._cfg.sample_rate, + channels=self._cfg.channels, + padding_ms=input_padding_ms, + ) + http_dur_ms = pcm_duration_ms(http_pcm, self._cfg.sample_rate, self._cfg.channels) + http_level_dbfs = dbfs_pcm16le(http_pcm) + if applied_input_padding_ms: + log_flow_event( + self._structured_logger or logger, + "stt_input_padding", + request_id=req_id, + padding_ms=applied_input_padding_ms, + original_audio_ms=round(dur_ms), + sent_audio_ms=round(http_dur_ms), + ) + if self._timeline is not None: + self._timeline.emit( + "stt_input_padding", + request_id=req_id, + padding_ms=applied_input_padding_ms, + original_audio_ms=round(dur_ms), + sent_audio_ms=round(http_dur_ms), + ) + + wav_buf = io.BytesIO() + with wave.open(wav_buf, "wb") as wf: + wf.setnchannels(self._cfg.channels) + wf.setsampwidth(2) + wf.setframerate(self._cfg.sample_rate) + wf.writeframes(http_pcm) + wav_bytes = wav_buf.getvalue() + + _dump_audio_for_debug( + req_id=req_id, + pcm=http_pcm, + wav_bytes=wav_bytes, + sample_rate=self._cfg.sample_rate, + channels=self._cfg.channels, + ) + + message_id = next_turn_message_id(self._structured_log_context) + enqueue_stt_vad_audio_upload( + wav_bytes=wav_bytes, + req_id=req_id, + message_id=message_id, + structured_log_context=self._structured_log_context, + metadata={ + "sample_rate": self._cfg.sample_rate, + "channels": self._cfg.channels, + "original_audio_ms": round(dur_ms), + "sent_audio_ms": round(http_dur_ms), + "original_dbfs": round(level_dbfs, 2), + "sent_dbfs": round(http_level_dbfs, 2), + "input_padding_ms": applied_input_padding_ms, + }, + logger_override=self._structured_logger or logger, + timeline=self._timeline, + ) + + # fallback: API interna + return await self._post_internal_http( + http_pcm, + wav_bytes, + req_id=req_id, + started_ns=started_ns, + language=language, + message_id=message_id, + audio_duration_ms=http_dur_ms, + original_audio_duration_ms=dur_ms, + input_padding_ms=applied_input_padding_ms, + level_dbfs=http_level_dbfs, + ) diff --git a/src/app/providers/stt_vosk.py b/src/app/providers/stt_vosk.py new file mode 100644 index 0000000..9c4b3b0 --- /dev/null +++ b/src/app/providers/stt_vosk.py @@ -0,0 +1,286 @@ +import json +from pathlib import Path +from typing import Optional, Dict, Any, Tuple, List + +import numpy as np +from vosk import Model, KaldiRecognizer, SetLogLevel +SetLogLevel(-1) + +DEFAULT_GRAMMAR = json.dumps([ + "sim", + "não", + "ok", + "cancela", + "cancelar", + "ligar", + "desligar", + + # +150 variações + "claro", + "claro que sim", + "com certeza", + "isso", + "isso mesmo", + "isso aí", + "isso ai", + "exato", + "exatamente", + "perfeito", + "beleza", + "certo", + "certinho", + "tá", + "ta", + "tá certo", + "ta certo", + "tá bom", + "ta bom", + "tudo bem", + "pode", + "pode sim", + "pode ser", + "pode ser sim", + "pode deixar", + "deixa", + "de acordo", + "concordo", + "concordo sim", + "confirmo", + "eu confirmo", + "confirmado", + "positivo", + "quero", + "manda ver", + "interessante", + "legal", + "bom", + + "negativo", + "de jeito nenhum", + "de forma alguma", + "nem pensar", + "não quero", + "não mesmo", + "não pode", + "não dá", + "não dá não", + "não quero não", + "quero não", + "não quero isso", + "não precisa", + "não preciso", + "dispenso", + "para", + "pare", + "parar", + "tá bom do jeito que tá", + + "tá ok", + "ta ok", + "tudo certo", + "tudo ok", + "certo então", + "beleza então", + "entendi", + "entendido", + "entendi sim", + "entendido sim", + "fechou", + "fechado", + "combinado", + "show", + "show de bola", + "top", + "maravilha", + "tranquilo", + "tranquila", + "tá tranquilo", + "ta tranquilo", + + "cancelamento", + "cancele", + "cancela aí", + "cancela ai", + "cancela agora", + "cancelar agora", + "pode cancelar", + "pode cancelar aí", + "pode cancelar ai", + "quero cancelar", + "eu quero cancelar", + "quero o cancelamento", + "faz o cancelamento", + "fazer cancelamento", + "solicitar cancelamento", + "desisto", + "desistir", + "quero desistir", + "para tudo", + "pare tudo", + "interrompe", + "interromper", + "encerra", + "encerrar", + "finaliza", + "caro", + + "liga", + "ligue", + "liga aí", + "liga ai", + "liga agora", + "ligar agora", + "pode ligar", + "pode ligar aí", + "pode ligar ai", + "pode ligar agora", + "conecta", + "conectar", + "conecte", + "inicia", + "iniciar", + "inicie", + "ativar", + "ativa", + "ative", + "habilitar", + + "desliga", + "desligue", + "desliga aí", + "desliga ai", + "desliga agora", + "desligar agora", + "pode desligar", + "pode desligar aí", + "pode desligar ai", + "pode desligar agora", + "desconecta", + "desconectar", + "desconecte", + "encerra a ligação", + "finaliza a ligação", + "desativar", + "desativa", + "desative" +], ensure_ascii=False) + + +class VoskSTT: + """ + STT com Vosk para CPU. + + Retorna SEMPRE: + {"data": {"text": "...", "words": [{"word","start","end","probability"}]}} + + Observação: + - Vosk retorna confiança por palavra em "conf" dentro de "result". + - A gente converte "conf" -> "probability". + """ + + def __init__( + self, + model_path: str, + sample_rate: int = 16000, + grammar_json: Optional[str] = DEFAULT_GRAMMAR, + ): + mp = Path(model_path) + if not mp.exists(): + raise FileNotFoundError(f"Vosk model_path não encontrado: {model_path}") + + self.sample_rate = sample_rate + self.grammar_json = grammar_json + self.model = Model(str(mp)) + + def transcribe_pcm16le(self, pcm_bytes: bytes) -> Dict[str, Any]: + """ + Recebe bytes PCM16LE mono (16k) e retorna dict no formato padrão. + """ + if not pcm_bytes: + return {"data": {"text": "", "words": []}} + + pcm16 = np.frombuffer(pcm_bytes, dtype=np.int16) # zero-copy + return self.recognize_pcm16(pcm16) + + def recognize_pcm16(self, pcm16: np.ndarray) -> Dict[str, Any]: + pcm16 = np.asarray(pcm16) + + if pcm16.dtype != np.int16: + raise TypeError(f"pcm16 must be int16, got {pcm16.dtype}") + + if pcm16.ndim == 2: + if pcm16.shape[1] != 1: + raise ValueError("Only mono supported; pass a single channel.") + pcm16 = pcm16[:, 0] + + # usa grammar se configurado + if self.grammar_json: + rec = KaldiRecognizer(self.model, self.sample_rate, self.grammar_json) + else: + rec = KaldiRecognizer(self.model, self.sample_rate) + + rec.SetWords(True) + rec.AcceptWaveform(pcm16.tobytes()) + raw = json.loads(rec.FinalResult()) + + return self.to_data_words_format(raw) + + @staticmethod + def to_data_words_format(raw: Dict[str, Any]) -> Dict[str, Any]: + """ + Converte o retorno do Vosk (raw) para: + {"data": {"text": "...", "words": [...]}} + """ + if not isinstance(raw, dict): + return {"data": {"text": str(raw), "words": []}} + + text = (raw.get("text") or "").strip() + words_out: List[Dict[str, Any]] = [] + + # formato padrão do Vosk com SetWords(True) + # {"text":"...", "result":[{"word","start","end","conf"}]} + if isinstance(raw.get("result"), list): + for w in raw["result"]: + if not isinstance(w, dict): + continue + words_out.append({ + "word": w.get("word", ""), + "start": float(w.get("start", 0.0) or 0.0), + "end": float(w.get("end", 0.0) or 0.0), + "probability": float(w.get("conf", w.get("probability", 0.0)) or 0.0), + }) + # fallback: se vier "words" já no formato desejado + elif isinstance(raw.get("words"), list): + for w in raw["words"]: + if not isinstance(w, dict): + continue + words_out.append({ + "word": w.get("word", ""), + "start": float(w.get("start", 0.0) or 0.0), + "end": float(w.get("end", 0.0) or 0.0), + "probability": float(w.get("probability", 0.0) or 0.0), + }) + + return {"data": {"text": text, "words": words_out}} + + @staticmethod + def result_to_text(res: Dict[str, Any]) -> str: + """ + Retorna apenas o texto do payload padronizado (ou do raw). + """ + if not isinstance(res, dict): + return str(res) + if isinstance(res.get("data"), dict): + t = res["data"].get("text") + return t.strip() if isinstance(t, str) else "" + t = res.get("text") + return t.strip() if isinstance(t, str) else "" + + @staticmethod + def result_to_text_and_payload(res: Dict[str, Any]) -> Tuple[str, Dict[str, Any]]: + """ + Retorna (texto, payload_padronizado). + """ + payload = VoskSTT.to_data_words_format(res) if not (isinstance(res, dict) and isinstance(res.get("data"), dict)) else res + text = VoskSTT.result_to_text(payload) + return text, payload diff --git a/src/app/providers/stt_word_dictionary.py b/src/app/providers/stt_word_dictionary.py new file mode 100644 index 0000000..69d59ae --- /dev/null +++ b/src/app/providers/stt_word_dictionary.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +SHORT_WORD_ALLOWLIST = { + "sim", + "não", + "nao", + "ok", + "tá", + "ta", + "certo", + "isso", + "confirmo", + "positivo", + "negativo", +} + +SHORT_WORD_ALLOWLIST_MIN_PROB = 0.0 diff --git a/src/app/providers/tts.py b/src/app/providers/tts.py new file mode 100644 index 0000000..33dbfc3 --- /dev/null +++ b/src/app/providers/tts.py @@ -0,0 +1,324 @@ +# app/providers/tts.py +from __future__ import annotations + +import os +import math +import struct +from math import ceil +from typing import TYPE_CHECKING, Any, Generator, List, Optional, Tuple + +from app.common.timed import timed + +if TYPE_CHECKING: + from elevenlabs import VoiceSettings + +SAMPLE_RATE = 16_000 +BYTES_PER_SAMPLE = 2 # PCM16 +CHANNELS = 1 + + +def frame_bytes_for_ms( + ms: int, + sr: int = SAMPLE_RATE, + ch: int = CHANNELS, + bps: int = BYTES_PER_SAMPLE, +) -> int: + if ms <= 0: + raise ValueError("frame_ms must be > 0") + samples_per_channel = round(sr * (ms / 1000.0)) + return int(samples_per_channel * ch * bps) + + +def _iter_pcm_frames( + pcm: bytes, + *, + frame_ms: int, + strict_frame_size: bool, +) -> Generator[bytes, None, None]: + if not pcm: + return + yield # pragma: no cover + + if frame_ms % 10 != 0: + raise ValueError("frame_ms should be a multiple of 10ms for telephony pipelines") + + frame_bytes = frame_bytes_for_ms(frame_ms) + offset = 0 + pcm_len = len(pcm) + while offset + frame_bytes <= pcm_len: + yield pcm[offset:offset + frame_bytes] + offset += frame_bytes + + if offset >= pcm_len: + return + + tail = pcm[offset:] + if strict_frame_size: + pad = frame_bytes - len(tail) + if pad > 0: + tail = tail + (b"\x00" * pad) + else: + tail = tail[:frame_bytes] + yield tail + + +def _frames_with_optional_silence( + *, + pcm_frames: Generator[bytes, None, None], + frame_ms: int, + pre_silence_ms: int, + post_silence_ms: int, +) -> List[bytes]: + frame_bytes = frame_bytes_for_ms(frame_ms) + silence_frame = b"\x00" * frame_bytes + frames: List[bytes] = [] + + if pre_silence_ms > 0: + frames.extend([silence_frame] * ceil(pre_silence_ms / frame_ms)) + + for frame in pcm_frames: + if len(frame) != frame_bytes: + if len(frame) < frame_bytes: + frame = frame + (b"\x00" * (frame_bytes - len(frame))) + else: + frame = frame[:frame_bytes] + frames.append(frame) + + if post_silence_ms > 0: + frames.extend([silence_frame] * ceil(post_silence_ms / frame_ms)) + + return frames + + +class TTS: + """ + ElevenLabs TTS -> PCM16 16kHz (mono). + """ + + def __init__( + self, + api_key: str, + voice_id: str, + model_id: str, + *, + voice_settings: Optional["VoiceSettings"] = None, + ) -> None: + if not api_key: + raise ValueError("api_key is required") + if not voice_id: + raise ValueError("voice_id is required") + + elevenlabs_sdk = _load_elevenlabs_sdk() + self.client = elevenlabs_sdk["client_cls"](api_key=api_key) + self.voice_id = voice_id + self.model_id = model_id + self.voice_settings = voice_settings or elevenlabs_sdk["voice_settings_cls"]( + stability=0.45, + speed=1.02, + similarity_boost=0.75, + use_speaker_boost=True, + style=0.7, + ) + + @timed("TTS") + def synthesize_pcm16k(self, text: str) -> bytes: + text = (text or "").strip() + if not text: + return b"" + + stream = self.client.text_to_speech.convert( + voice_id=self.voice_id, + output_format="pcm_16000", + text=text, + model_id=self.model_id, + voice_settings=self.voice_settings, + ) + + out = bytearray() + for chunk in stream: + if chunk: + out.extend(chunk) + return bytes(out) + + @timed("TTS") + def synthesize_frames( + self, + text: str, + frame_ms: int = 20, + *, + strict_frame_size: bool = True, + ) -> Generator[bytes, None, None]: + yield from _iter_pcm_frames( + self.synthesize_pcm16k(text), + frame_ms=frame_ms, + strict_frame_size=strict_frame_size, + ) + + @timed("Intro TTS") + def synthesize_ws_frames( + self, + text: str, + *, + frame_ms: int = 20, + pre_silence_ms: int = 0, + post_silence_ms: int = 0, + ) -> List[bytes]: + text = (text or "").strip() + if not text: + return [] + + return _frames_with_optional_silence( + pcm_frames=self.synthesize_frames(text, frame_ms=frame_ms, strict_frame_size=True), + frame_ms=frame_ms, + pre_silence_ms=pre_silence_ms, + post_silence_ms=post_silence_ms, + ) + + +class FakeTTS: + """ + Provider local para smoke tests. + + Gera um tom PCM16 simples com duracao proporcional ao texto para validar o + pipeline sem depender de servicos externos de TTS. + """ + + def __init__( + self, + *, + tone_hz: int = 440, + amplitude: float = 0.12, + min_duration_ms: int = 320, + max_duration_ms: int = 2200, + char_duration_ms: int = 35, + ) -> None: + self.tone_hz = max(120, int(tone_hz)) + self.amplitude = min(max(float(amplitude), 0.0), 0.95) + self.min_duration_ms = max(80, int(min_duration_ms)) + self.max_duration_ms = max(self.min_duration_ms, int(max_duration_ms)) + self.char_duration_ms = max(5, int(char_duration_ms)) + + def _duration_ms_for_text(self, text: str) -> int: + text = (text or "").strip() + if not text: + return 0 + + estimated = max(self.min_duration_ms, len(text) * self.char_duration_ms) + return min(estimated, self.max_duration_ms) + + def _envelope(self, index: int, total_samples: int) -> float: + fade_samples = max(1, int(SAMPLE_RATE * 0.02)) + if index < fade_samples: + return index / fade_samples + tail_start = max(0, total_samples - fade_samples) + if index >= tail_start: + return max(0.0, (total_samples - index) / fade_samples) + return 1.0 + + @timed("TTS") + def synthesize_pcm16k(self, text: str) -> bytes: + duration_ms = self._duration_ms_for_text(text) + if duration_ms <= 0: + return b"" + + total_samples = max(1, int(SAMPLE_RATE * duration_ms / 1000)) + out = bytearray() + angular = (2.0 * math.pi * self.tone_hz) / SAMPLE_RATE + + for sample_index in range(total_samples): + env = self._envelope(sample_index, total_samples) + value = math.sin(sample_index * angular) * self.amplitude * env + out.extend(struct.pack(" Generator[bytes, None, None]: + yield from _iter_pcm_frames( + self.synthesize_pcm16k(text), + frame_ms=frame_ms, + strict_frame_size=strict_frame_size, + ) + + @timed("Intro TTS") + def synthesize_ws_frames( + self, + text: str, + *, + frame_ms: int = 20, + pre_silence_ms: int = 0, + post_silence_ms: int = 0, + ) -> List[bytes]: + text = (text or "").strip() + if not text: + return [] + + return _frames_with_optional_silence( + pcm_frames=self.synthesize_frames(text, frame_ms=frame_ms, strict_frame_size=True), + frame_ms=frame_ms, + pre_silence_ms=pre_silence_ms, + post_silence_ms=post_silence_ms, + ) + + +def build_tts_provider_from_env( + provider: str, + *, + voice_id: str = "", + model_id: str = "", + language: str = "", +) -> Tuple[object | None, str | None]: + normalized = (provider or "elevenlabs").strip().lower() + + if normalized == "fake": + return ( + FakeTTS( + tone_hz=int(os.getenv("FAKE_TTS_TONE_HZ", "440") or "440"), + amplitude=float(os.getenv("FAKE_TTS_AMPLITUDE", "0.12") or "0.12"), + min_duration_ms=int(os.getenv("FAKE_TTS_MIN_DURATION_MS", "320") or "320"), + max_duration_ms=int(os.getenv("FAKE_TTS_MAX_DURATION_MS", "2200") or "2200"), + char_duration_ms=int(os.getenv("FAKE_TTS_CHAR_DURATION_MS", "35") or "35"), + ), + None, + ) + + if normalized in {"", "elevenlabs"}: + api_key = (os.getenv("ELEVENLABS_API_KEY", "") or "").strip() + resolved_voice_id = (voice_id or os.getenv("ELEVENLABS_VOICE_ID", "")).strip() + resolved_model_id = (model_id or os.getenv("ELEVENLABS_MODEL_ID", "eleven_flash_v2_5")).strip() + if not api_key or not resolved_voice_id: + return None, "missing_elevenlabs_config" + if not _is_elevenlabs_available(): + return None, "missing_elevenlabs_sdk" + return TTS(api_key=api_key, voice_id=resolved_voice_id, model_id=resolved_model_id), None + + return None, f"unsupported_tts_provider:{normalized}" + + +def _load_elevenlabs_sdk() -> dict[str, Any]: + try: + from elevenlabs import VoiceSettings + from elevenlabs.client import ElevenLabs + except ModuleNotFoundError as exc: + raise RuntimeError( + "ElevenLabs SDK is not installed. Install dependency `elevenlabs` to enable this provider." + ) from exc + + return { + "voice_settings_cls": VoiceSettings, + "client_cls": ElevenLabs, + } + + +def _is_elevenlabs_available() -> bool: + try: + _load_elevenlabs_sdk() + except RuntimeError: + return False + return True diff --git a/src/app/services/session_context.py b/src/app/services/session_context.py new file mode 100644 index 0000000..6cc1f95 --- /dev/null +++ b/src/app/services/session_context.py @@ -0,0 +1,58 @@ +from __future__ import annotations + +from typing import Any, Mapping + + +def _pick_value(payload: Mapping[str, Any] | None, *keys: str) -> str: + if not isinstance(payload, Mapping): + return "" + + for key in keys: + value = payload.get(key) + if value not in (None, ""): + return str(value).strip() + + lowered = {str(key).lower(): value for key, value in payload.items()} + for key in keys: + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return str(value).strip() + + return "" + + +def extract_protocol( + start_payload: Mapping[str, Any] | None, + session_data: Mapping[str, Any] | None = None, +) -> str: + data = start_payload.get("data") if isinstance(start_payload, Mapping) else {} + if not isinstance(data, Mapping): + data = {} + + protocol = _pick_value( + data, + "protocol_id", + "protocolId", + "protocolo", + "PROTOCOLO", + "protocol", + "Protocol", + "routerCallKey", + "RouterCallKey", + "router_call_key", + ) + if protocol: + return protocol + + return _pick_value( + session_data, + "protocol_id", + "protocolId", + "protocolo", + "PROTOCOLO", + "protocol", + "Protocol", + "routerCallKey", + "RouterCallKey", + "router_call_key", + ) diff --git a/src/app/services/text_pipeline.py b/src/app/services/text_pipeline.py new file mode 100644 index 0000000..6339db2 --- /dev/null +++ b/src/app/services/text_pipeline.py @@ -0,0 +1,162 @@ +import asyncio +import json +import logging +import time + +from fastapi import WebSocket, WebSocketDisconnect + +from agent.pipeline.customer_pipeline_langgraph import CustomerPipeline + +from app.services.session_context import extract_protocol + +logger = logging.getLogger(__name__) + + +class TextSessionPipeline: + """ + Pipeline puramente texto: + - envia "ready" + - envia a abertura do agente via agent.start() + - para cada mensagem de texto recebida, responde com agent.run(text) + Protocolo sugerido (JSON): + -> {"type":"start", "data": {...}} # primeiro pacote + -> {"type":"text", "text":"..."} # mensagens do usuário + -> {"type":"stop"} # encerra + Eventos enviados pelo servidor: + <- {"type":"ready", "session_id": "..."} + <- {"type":"agent_start", "text": "..."} + <- {"type":"agent_reply", "text": "..."} + <- {"type":"bye"} + <- {"type":"error","message":"..."} + """ + def __init__( + self, + ws: WebSocket, + start_payload: dict, + session_id: str, + *, + session_data: dict, + intro: str, + ): + self.ws = ws + self.start_payload = start_payload or {} + self.session_id = session_id + self.session_data = dict(session_data or {}) + self.intro = (intro or "").strip() + self.protocol = extract_protocol(self.start_payload, self.session_data) + self.agent = CustomerPipeline(self.session_data) + self._close = False + + async def send_json(self, payload: dict): + await self.ws.send_text(json.dumps(payload, ensure_ascii=False)) + + async def run(self): + # 1) sinaliza pronto + await self.send_json({"type": "ready", "session_id": self.session_id}) + + # 2) mensagem inicial do agente + stage = "" + try: + stage, opening = self.agent.start() + except Exception as e: + logger.exception("Erro no agent.start()") + await self.send_json({"type": "error", "message": f"start_failed: {e}"}) + opening = None + + if not opening: + opening = self.intro or None + + if opening: + logger.info( + "TEXT_AGENT_START | session_id=%s | stage=%s | text=%r", + self.session_id, + stage, + opening, + ) + await self.send_json({"type": "agent_start", "text": opening}) + + self.agent.prepare(True, self.protocol) + # 3) loop de interação + while not self._close: + try: + msg = await asyncio.wait_for(self.ws.receive(), timeout=300) + except asyncio.TimeoutError: + # mantém conexão viva (opcional) + try: + await self.ws.send_text("keepalive") + except Exception: + pass + continue + except WebSocketDisconnect: + logger.info(f"Cliente desconectou: {self.session_id}") + break + except Exception as e: + logger.exception("Erro no loop text ws") + await self.send_json({"type":"error", "message": str(e)}) + break + + if msg.get("type") == "websocket.disconnect": + break + + if msg.get("type") != "websocket.receive": + continue + + # Só tratamos conteúdo textual nessa rota + raw_text = msg.get("text") + if not raw_text: + # ignora binários nesta rota + continue + + # Aceita texto puro ou JSON + try: + data = json.loads(raw_text) + except Exception: + data = {"type": "text", "text": raw_text} + + mtype = data.get("type") + if mtype in ("stop", "bye"): + await self.send_json({"type": "bye"}) + break + + if mtype in ("text", "user_input", "message"): + text = (data.get("text") or "").strip() + if not text: + continue + logger.info( + "TEXT_USER | session_id=%s | protocol=%s | text=%r", + self.session_id, + self.protocol, + text, + ) + + try: + agent_started = time.perf_counter() + stage, reply = self.agent.run(text) + agent_duration_ms = round((time.perf_counter() - agent_started) * 1000) + self.agent.set_interruption(was_interrupted=False) + #print(f"stage <> {stage}") + self._close = str(stage) in ("Concluído", "DONE", "DONE_CONCLUIDO") + + + except Exception as e: + logger.exception("Erro no agent.run()") + await self.send_json({"type": "error", "message": f"run_failed: {e}"}) + continue + + logger.info( + "TEXT_AGENT_REPLY | session_id=%s | stage=%s | close=%s | duration_ms=%s | text=%r", + self.session_id, + stage, + self._close, + agent_duration_ms, + reply, + ) + await self.send_json({"type": "agent_reply", "text": reply}) + if(self._close): break + else: + await self.send_json({"type": "error", "message": f"unknown_type:{mtype}"}) + try: + await self.ws.close() + output = await self.agent._end_service() + except Exception: + logger.exception("Erro ao encerrar pipeline de texto") diff --git a/src/app/services/text_pipeline_stream.py b/src/app/services/text_pipeline_stream.py new file mode 100644 index 0000000..6aa148f --- /dev/null +++ b/src/app/services/text_pipeline_stream.py @@ -0,0 +1,195 @@ +import asyncio +import json +import logging +import time + +from fastapi import WebSocket, WebSocketDisconnect + +from agent.pipeline.pipeline_streaming_text import CustomerPipeline +from app.services.session_context import extract_protocol + +logger = logging.getLogger(__name__) + + +class TextSessionPipelineStream: + """ + Pipeline para comunicação texto via WebSocket com streaming token a token. + + Protocolo sugerido (JSON): + -> {"type":"start", "data": {...}} + -> {"type":"text", "text":"..."} + -> {"type":"stop"} + + Eventos enviados pelo servidor: + <- {"type":"ready", "session_id": "..."} + <- {"type":"agent_start", "text": "..."} + <- {"type":"agent_reply_chunk", "delta": "token"} + <- {"type":"agent_reply_end", "text": "resposta completa"} + <- {"type":"bye"} + <- {"type":"error","message":"..."} + """ + + def __init__( + self, + ws: WebSocket, + start_payload: dict, + session_id: str, + *, + session_data: dict | None = None, + intro: str = "", + ): + self.ws = ws + self.start_payload = start_payload or {} + self.session_id = session_id + self.session_data = dict(session_data or {}) + self.intro = (intro or "").strip() + self.protocol = extract_protocol(self.start_payload, self.session_data) + self.agent = CustomerPipeline() + self._close = False + + async def send_json(self, payload: dict): + await self.ws.send_text(json.dumps(payload, ensure_ascii=False)) + + # ============================================================ + # STREAMING TOKEN-A-TOKEN + # ============================================================ + + async def _to_async_gen(self, generator): + """Transforma generator síncrono em assíncrono.""" + for item in generator: + yield item + + async def stream_reply(self, text: str): + """ + Envia resposta token-a-token via WebSocket: + - {"type":"agent_reply_chunk", "delta": token} + - {"type":"agent_reply_end", "text": full_response} + """ + try: + agent_started = time.perf_counter() + token_generator = self.agent.run(text) + except Exception as e: + await self.send_json({"type": "error", "message": f"stream_init_failed: {e}"}) + return + + full_text = "" + + try: + async for token in self._to_async_gen(token_generator): + full_text += token + await self.send_json({"type": "agent_reply_chunk", "delta": token}) + await asyncio.sleep(0) + except Exception as e: + await self.send_json({"type": "error", "message": f"stream_failed: {e}"}) + return + + # Finaliza streaming + agent_duration_ms = round((time.perf_counter() - agent_started) * 1000) + logger.info( + "TEXT_STREAM_AGENT_REPLY | session_id=%s | duration_ms=%s | text=%r", + self.session_id, + agent_duration_ms, + full_text, + ) + await self.send_json({ + "type": "agent_reply_end", + "text": full_text + }) + + # ============================================================ + + async def run(self): + # 1) envia ready + await self.send_json({"type": "ready", "session_id": self.session_id}) + + # 2) mensagem inicial do agente + stage = "" + try: + stage, opening = self.agent.start() + except Exception as e: + logger.exception("Erro no agent.start()") + await self.send_json({"type": "error", "message": f"start_failed: {e}"}) + opening = None + + if not opening: + opening = self.intro or None + + if opening: + logger.info( + "TEXT_STREAM_AGENT_START | session_id=%s | stage=%s | text=%r", + self.session_id, + stage, + opening, + ) + await self.send_json({"type": "agent_start", "text": opening}) + + self.agent.prepare(True, self.protocol) + + # 3) loop WS + while not self._close: + try: + msg = await asyncio.wait_for(self.ws.receive(), timeout=300) + except asyncio.TimeoutError: + # ping opcional + try: + await self.ws.send_text("keepalive") + except Exception: + pass + continue + except WebSocketDisconnect: + logger.info(f"Cliente desconectou: {self.session_id}") + break + except Exception as e: + logger.exception("Erro no loop text ws") + await self.send_json({"type": "error", "message": str(e)}) + break + + if msg.get("type") == "websocket.disconnect": + break + + if msg.get("type") != "websocket.receive": + continue + + raw_text = msg.get("text") + if not raw_text: + continue + + # aceita JSON ou texto puro + try: + data = json.loads(raw_text) + except Exception: + data = {"type": "text", "text": raw_text} + + mtype = data.get("type") + + if mtype in ("stop", "bye"): + await self.send_json({"type": "bye"}) + break + + if mtype in ("text", "user_input", "message"): + text = (data.get("text") or "").strip() + if not text: + continue + logger.info( + "TEXT_STREAM_USER | session_id=%s | protocol=%s | text=%r", + self.session_id, + self.protocol, + text, + ) + + # STREAMING TOKEN A TOKEN + await self.stream_reply(text) + + # depois do streaming, atualiza stage + try: + self._close = str(self.agent.stage) in ("Concluído", "DONE", "DONE_CONCLUIDO") + except Exception as e: + await self.send_json({"type": "error", "message": f"run_failed: {e}"}) + + else: + await self.send_json({"type": "error", "message": f"unknown_type:{mtype}"}) + + try: + await self.ws.close() + except Exception: + pass diff --git a/src/app/tools/__init__.py b/src/app/tools/__init__.py new file mode 100644 index 0000000..80fe0d2 --- /dev/null +++ b/src/app/tools/__init__.py @@ -0,0 +1 @@ +"""Operational tooling entrypoints for the TIA application.""" diff --git a/src/app/tools/local_stresstest/__init__.py b/src/app/tools/local_stresstest/__init__.py new file mode 100644 index 0000000..af61387 --- /dev/null +++ b/src/app/tools/local_stresstest/__init__.py @@ -0,0 +1,10 @@ +"""Local operational stress test for STT, TTS, Bridge, and LiveKit.""" + +from app.tools.local_stresstest.text import TextComparison, compare_text, normalize_text, word_error_rate + +__all__ = [ + "TextComparison", + "compare_text", + "normalize_text", + "word_error_rate", +] diff --git a/src/app/tools/local_stresstest/__main__.py b/src/app/tools/local_stresstest/__main__.py new file mode 100644 index 0000000..93ee8ee --- /dev/null +++ b/src/app/tools/local_stresstest/__main__.py @@ -0,0 +1,13 @@ +from __future__ import annotations + +import asyncio + +from app.tools.local_stresstest.runner import run + + +def main() -> None: + raise SystemExit(asyncio.run(run())) + + +if __name__ == "__main__": + main() diff --git a/src/app/tools/local_stresstest/audio.py b/src/app/tools/local_stresstest/audio.py new file mode 100644 index 0000000..8c4728e --- /dev/null +++ b/src/app/tools/local_stresstest/audio.py @@ -0,0 +1,364 @@ +from __future__ import annotations + +import audioop +import math +import random +import wave +from dataclasses import dataclass +from pathlib import Path + + +TARGET_SAMPLE_RATE = 16_000 +TARGET_CHANNELS = 1 +SAMPLE_WIDTH = 2 +FRAME_MS = 20 + + +@dataclass(frozen=True, slots=True) +class AudioSample: + name: str + pcm: bytes + sample_rate: int = TARGET_SAMPLE_RATE + channels: int = TARGET_CHANNELS + + +@dataclass(frozen=True, slots=True) +class AudioMetrics: + duration_ms: int + rms_dbfs: float + peak_dbfs: float + clipping_ratio: float + bytes_len: int + + @property + def audible(self) -> bool: + return self.rms_dbfs > -55.0 + + +@dataclass(frozen=True, slots=True) +class VADProxyMetrics: + threshold_dbfs: float + prefix_padding_ms: int + min_speech_ms: int + first_voice_ms: int + last_voice_ms: int + speech_ms: int + padded_start_ms: int + unrecovered_prefix_ms: int + initial_200ms_dbfs: float + initial_500ms_dbfs: float + low_start_risk: bool + + +def _dbfs(value: float) -> float: + if value <= 0: + return -120.0 + return 20.0 * math.log10(value / 32768.0) + + +def _samples_from_pcm(pcm: bytes) -> list[int]: + if not pcm: + return [] + count = len(pcm) // SAMPLE_WIDTH + return list(audioop.getsample(pcm, SAMPLE_WIDTH, idx) for idx in range(count)) + + +def _pcm_from_samples(samples: list[int]) -> bytes: + out = bytearray() + for value in samples: + clipped = max(-32768, min(32767, int(round(value)))) + out.extend(clipped.to_bytes(2, byteorder="little", signed=True)) + return bytes(out) + + +def _silence_pcm(sample_rate: int, duration_ms: int) -> bytes: + return b"\x00" * round(sample_rate * duration_ms / 1000) * SAMPLE_WIDTH + + +def _split_pcm_at_ms(pcm: bytes, *, sample_rate: int, split_ms: int) -> tuple[bytes, bytes]: + byte_count = max(0, round(sample_rate * split_ms / 1000) * SAMPLE_WIDTH) + byte_count -= byte_count % SAMPLE_WIDTH + return pcm[:byte_count], pcm[byte_count:] + + +def audio_metrics(pcm: bytes, *, sample_rate: int = TARGET_SAMPLE_RATE, channels: int = TARGET_CHANNELS) -> AudioMetrics: + samples = _samples_from_pcm(pcm) + if not samples: + return AudioMetrics(duration_ms=0, rms_dbfs=-120.0, peak_dbfs=-120.0, clipping_ratio=0.0, bytes_len=0) + + rms = audioop.rms(pcm, SAMPLE_WIDTH) + peak = max(abs(item) for item in samples) + clipped = sum(1 for item in samples if abs(item) >= 32760) + sample_count_per_channel = len(samples) / max(1, channels) + duration_ms = round((sample_count_per_channel / max(1, sample_rate)) * 1000) + return AudioMetrics( + duration_ms=duration_ms, + rms_dbfs=round(_dbfs(rms), 2), + peak_dbfs=round(_dbfs(peak), 2), + clipping_ratio=round(clipped / len(samples), 6), + bytes_len=len(pcm), + ) + + +def _window_dbfs( + pcm: bytes, + *, + sample_rate: int = TARGET_SAMPLE_RATE, + duration_ms: int, +) -> float: + byte_count = max(0, round(sample_rate * duration_ms / 1000) * SAMPLE_WIDTH) + byte_count -= byte_count % SAMPLE_WIDTH + window = pcm[:byte_count] if byte_count else b"" + if not window: + return -120.0 + return round(_dbfs(audioop.rms(window, SAMPLE_WIDTH)), 2) + + +def vad_proxy_metrics( + sample: AudioSample, + *, + threshold_dbfs: float = -45.0, + prefix_padding_ms: int = 1000, + min_speech_ms: int = 100, + frame_ms: int = FRAME_MS, +) -> VADProxyMetrics: + frame_bytes = max(1, round(sample.sample_rate * frame_ms / 1000)) * SAMPLE_WIDTH + threshold_rms = int(32768 * (10 ** (threshold_dbfs / 20.0))) + active_frames: list[int] = [] + total_frames = max(1, math.ceil(len(sample.pcm) / frame_bytes)) + for frame_idx in range(total_frames): + frame = sample.pcm[frame_idx * frame_bytes : (frame_idx + 1) * frame_bytes] + if len(frame) < frame_bytes: + frame = frame + b"\x00" * (frame_bytes - len(frame)) + if audioop.rms(frame, SAMPLE_WIDTH) >= threshold_rms: + active_frames.append(frame_idx) + + initial_200ms_dbfs = _window_dbfs(sample.pcm, sample_rate=sample.sample_rate, duration_ms=200) + initial_500ms_dbfs = _window_dbfs(sample.pcm, sample_rate=sample.sample_rate, duration_ms=500) + if not active_frames: + return VADProxyMetrics( + threshold_dbfs=threshold_dbfs, + prefix_padding_ms=prefix_padding_ms, + min_speech_ms=min_speech_ms, + first_voice_ms=-1, + last_voice_ms=-1, + speech_ms=0, + padded_start_ms=-1, + unrecovered_prefix_ms=-1, + initial_200ms_dbfs=initial_200ms_dbfs, + initial_500ms_dbfs=initial_500ms_dbfs, + low_start_risk=True, + ) + + first_voice_ms = active_frames[0] * frame_ms + last_voice_ms = (active_frames[-1] + 1) * frame_ms + speech_ms = max(0, last_voice_ms - first_voice_ms) + padded_start_ms = max(0, first_voice_ms - prefix_padding_ms) + unrecovered_prefix_ms = padded_start_ms + low_start_risk = ( + speech_ms < min_speech_ms + or first_voice_ms > 0 + or unrecovered_prefix_ms > 0 + or initial_200ms_dbfs < threshold_dbfs + ) + return VADProxyMetrics( + threshold_dbfs=threshold_dbfs, + prefix_padding_ms=prefix_padding_ms, + min_speech_ms=min_speech_ms, + first_voice_ms=first_voice_ms, + last_voice_ms=last_voice_ms, + speech_ms=speech_ms, + padded_start_ms=padded_start_ms, + unrecovered_prefix_ms=unrecovered_prefix_ms, + initial_200ms_dbfs=initial_200ms_dbfs, + initial_500ms_dbfs=initial_500ms_dbfs, + low_start_risk=low_start_risk, + ) + + +def read_wav_mono16(path: Path | str, *, target_sample_rate: int = TARGET_SAMPLE_RATE) -> AudioSample: + wav_path = Path(path) + with wave.open(str(wav_path), "rb") as handle: + channels = handle.getnchannels() + sample_width = handle.getsampwidth() + sample_rate = handle.getframerate() + pcm = handle.readframes(handle.getnframes()) + + if sample_width != SAMPLE_WIDTH: + raise ValueError(f"{wav_path} must be PCM16 WAV, got sample_width={sample_width}") + if channels != TARGET_CHANNELS: + pcm = audioop.tomono(pcm, SAMPLE_WIDTH, 0.5, 0.5) + channels = TARGET_CHANNELS + if sample_rate != target_sample_rate: + pcm, _ = audioop.ratecv(pcm, SAMPLE_WIDTH, channels, sample_rate, target_sample_rate, None) + sample_rate = target_sample_rate + + return AudioSample(name=wav_path.stem, pcm=pcm, sample_rate=sample_rate, channels=channels) + + +def write_wav(path: Path | str, sample: AudioSample | bytes, *, sample_rate: int = TARGET_SAMPLE_RATE) -> Path: + wav_path = Path(path) + wav_path.parent.mkdir(parents=True, exist_ok=True) + if isinstance(sample, AudioSample): + pcm = sample.pcm + sample_rate = sample.sample_rate + channels = sample.channels + else: + pcm = sample + channels = TARGET_CHANNELS + + with wave.open(str(wav_path), "wb") as handle: + handle.setnchannels(channels) + handle.setsampwidth(SAMPLE_WIDTH) + handle.setframerate(sample_rate) + handle.writeframes(pcm) + return wav_path + + +def pcm_to_wav_bytes(sample: AudioSample) -> bytes: + import io + + buffer = io.BytesIO() + with wave.open(buffer, "wb") as handle: + handle.setnchannels(sample.channels) + handle.setsampwidth(SAMPLE_WIDTH) + handle.setframerate(sample.sample_rate) + handle.writeframes(sample.pcm) + return buffer.getvalue() + + +def _mulaw_roundtrip(pcm: bytes) -> bytes: + mulaw = audioop.lin2ulaw(pcm, SAMPLE_WIDTH) + return audioop.ulaw2lin(mulaw, SAMPLE_WIDTH) + + +def _apply_gain(pcm: bytes, gain: float) -> bytes: + return audioop.mul(pcm, SAMPLE_WIDTH, gain) + + +def _add_gaussian_noise_to_pcm(pcm: bytes, *, snr_db: float, seed: int) -> bytes: + samples = _samples_from_pcm(pcm) + if not samples: + return pcm + + rng = random.Random(seed) + signal_rms = math.sqrt(sum(float(item) * item for item in samples) / len(samples)) + noise_rms = signal_rms / (10 ** (snr_db / 20.0)) if signal_rms > 0 else 0.0 + out = [] + for value in samples: + out.append(value + rng.gauss(0.0, noise_rms)) + return _pcm_from_samples(out) + + +def _noise_pcm(*, duration_ms: int, sample_rate: int, dbfs: float, seed: int) -> bytes: + rng = random.Random(seed) + amplitude = 32768 * (10 ** (dbfs / 20.0)) + samples = round(sample_rate * duration_ms / 1000) + return _pcm_from_samples([rng.gauss(0.0, amplitude) for _ in range(samples)]) + + +def _add_silence(sample: AudioSample, *, before_ms: int, after_ms: int) -> AudioSample: + before = _silence_pcm(sample.sample_rate, before_ms) + after = _silence_pcm(sample.sample_rate, after_ms) + return AudioSample(name="leading_trailing_silence", pcm=before + sample.pcm + after) + + +def _add_noise(sample: AudioSample, *, snr_db: float, seed: int) -> AudioSample: + return AudioSample( + name=f"noise_snr_{round(snr_db)}", + pcm=_add_gaussian_noise_to_pcm(sample.pcm, snr_db=snr_db, seed=seed), + ) + + +def _with_name(name: str, pcm: bytes, sample: AudioSample) -> AudioSample: + return AudioSample(name=name, pcm=pcm, sample_rate=sample.sample_rate, channels=sample.channels) + + +def _add_named_silence(sample: AudioSample, *, name: str, before_ms: int, after_ms: int) -> AudioSample: + return _with_name( + name, + _silence_pcm(sample.sample_rate, before_ms) + sample.pcm + _silence_pcm(sample.sample_rate, after_ms), + sample, + ) + + +def _add_prefix_noise(sample: AudioSample, *, name: str, before_ms: int, dbfs: float, seed: int) -> AudioSample: + return _with_name( + name, + _noise_pcm(duration_ms=before_ms, sample_rate=sample.sample_rate, dbfs=dbfs, seed=seed) + sample.pcm, + sample, + ) + + +def _apply_prefix_gain(sample: AudioSample, *, name: str, prefix_ms: int, gain: float) -> AudioSample: + prefix, rest = _split_pcm_at_ms(sample.pcm, sample_rate=sample.sample_rate, split_ms=prefix_ms) + return _with_name(name, _apply_gain(prefix, gain) + rest, sample) + + +def _fade_in(sample: AudioSample, *, name: str, duration_ms: int) -> AudioSample: + samples = _samples_from_pcm(sample.pcm) + fade_samples = max(1, round(sample.sample_rate * duration_ms / 1000)) + out = [] + for idx, value in enumerate(samples): + if idx < fade_samples: + gain = idx / fade_samples + out.append(value * gain) + else: + out.append(value) + return _with_name(name, _pcm_from_samples(out), sample) + + +def _prefix_gain_with_noise( + sample: AudioSample, + *, + name: str, + prefix_ms: int, + gain: float, + snr_db: float, + seed: int, +) -> AudioSample: + prefix, rest = _split_pcm_at_ms(sample.pcm, sample_rate=sample.sample_rate, split_ms=prefix_ms) + quiet_prefix = _apply_gain(prefix, gain) + noisy_prefix = _add_gaussian_noise_to_pcm(quiet_prefix, snr_db=snr_db, seed=seed) + return _with_name(name, noisy_prefix + rest, sample) + + +def build_variations(sample: AudioSample) -> list[AudioSample]: + return [ + AudioSample(name="clean", pcm=sample.pcm), + AudioSample(name="low_volume", pcm=_apply_gain(sample.pcm, 0.45)), + AudioSample(name="very_low_volume", pcm=_apply_gain(sample.pcm, 0.22)), + AudioSample(name="high_volume", pcm=_apply_gain(sample.pcm, 1.55)), + AudioSample(name="clipped_high_volume", pcm=_apply_gain(sample.pcm, 2.6)), + _add_silence(sample, before_ms=600, after_ms=900), + _add_named_silence(sample, name="short_leading_silence", before_ms=120, after_ms=400), + _add_named_silence(sample, name="long_leading_silence", before_ms=1600, after_ms=900), + _add_noise(sample, snr_db=20, seed=20), + _add_noise(sample, snr_db=15, seed=15), + _add_noise(sample, snr_db=10, seed=10), + _add_prefix_noise(sample, name="pre_noise_300ms", before_ms=300, dbfs=-34.0, seed=300), + _add_prefix_noise(sample, name="pre_noise_800ms", before_ms=800, dbfs=-34.0, seed=800), + AudioSample(name="telephony_profile", pcm=_mulaw_roundtrip(sample.pcm)), + AudioSample(name="telephony_low_volume", pcm=_apply_gain(_mulaw_roundtrip(sample.pcm), 0.45)), + _fade_in(sample, name="initial_fade_in_250ms", duration_ms=250), + _fade_in(sample, name="initial_fade_in_500ms", duration_ms=500), + _fade_in(sample, name="initial_fade_in_900ms", duration_ms=900), + _apply_prefix_gain(sample, name="initial_dip_300ms", prefix_ms=300, gain=0.20), + _apply_prefix_gain(sample, name="initial_dip_700ms", prefix_ms=700, gain=0.25), + _apply_prefix_gain(sample, name="prefix_300ms_15pct", prefix_ms=300, gain=0.15), + _apply_prefix_gain(sample, name="prefix_600ms_20pct", prefix_ms=600, gain=0.20), + _apply_prefix_gain(sample, name="prefix_900ms_25pct", prefix_ms=900, gain=0.25), + _prefix_gain_with_noise( + sample, + name="low_prefix_noise_600ms", + prefix_ms=600, + gain=0.22, + snr_db=8, + seed=608, + ), + _with_name( + "low_prefix_telephony_700ms", + _mulaw_roundtrip(_apply_prefix_gain(sample, name="unused", prefix_ms=700, gain=0.25).pcm), + sample, + ), + ] diff --git a/src/app/tools/local_stresstest/report.py b/src/app/tools/local_stresstest/report.py new file mode 100644 index 0000000..6ea34cf --- /dev/null +++ b/src/app/tools/local_stresstest/report.py @@ -0,0 +1,436 @@ +from __future__ import annotations + +import csv +import json +from dataclasses import asdict, is_dataclass +from html import escape +from pathlib import Path +from typing import Any + +from app.tools.local_stresstest.scenarios import scenario_description + + +def _svg_pipeline(title: str, steps: list[str]) -> str: + box_w = 150 + box_h = 56 + gap = 34 + margin_x = 34 + margin_y = 64 + width = margin_x * 2 + len(steps) * box_w + (len(steps) - 1) * gap + height = 170 + title_x = width // 2 + parts = [ + f'', + "", + '', + '', + "", + "", + '', + f'{escape(title)}', + ] + y = margin_y + for idx, step in enumerate(steps): + x = margin_x + idx * (box_w + gap) + parts.append(f'') + words = step.split() + if len(words) > 2: + mid = (len(words) + 1) // 2 + lines = [" ".join(words[:mid]), " ".join(words[mid:])] + else: + lines = [step] + text_y = y + 28 - (len(lines) - 1) * 9 + for line in lines: + parts.append( + f'{escape(line)}' + ) + text_y += 18 + if idx < len(steps) - 1: + x1 = x + box_w + 6 + x2 = x + box_w + gap - 8 + y_mid = y + box_h / 2 + parts.append( + f'' + ) + parts.append("") + return "\n".join(parts) + "\n" + + +STT_MERMAID = "\n".join( + [ + "flowchart LR", + " A[WAV base] --> B[Variacao de audio]", + " B --> C[VAD proxy RMS]", + " C --> D[Sofya STT HTTP]", + " D --> E[Normalizar transcricao]", + " E --> F[Prefixo WER termos]", + " F --> G[Linha no CSV e Markdown]", + ] +) + +TTS_MERMAID = "\n".join( + [ + "flowchart LR", + " A[Texto de teste] --> B[xAI websocket TTS]", + " B --> C[PCM audio]", + " C --> D[Artefato WAV]", + " C --> E[Duracao, RMS e clipping]", + " E --> F[Linha no CSV e Markdown]", + ] +) + +E2E_MERMAID = "\n".join( + [ + "sequenceDiagram", + " participant Runner", + " participant Bridge", + " participant LiveKit", + " participant Agent", + " participant Sofya", + " participant XAI", + " Runner->>Bridge: start WS + frames PCM", + " Bridge->>LiveKit: publica audio do cliente", + " LiveKit->>Agent: track de audio", + " Agent->>Sofya: reconhece fala", + " Sofya-->>Agent: transcricao", + " Agent->>Agent: resposta remote_ws_fake", + " Agent->>XAI: sintetiza resposta", + " XAI-->>Agent: audio PCM", + " Agent-->>LiveKit: track de audio", + " LiveKit-->>Bridge: audio do agente", + " Bridge-->>Runner: frames PCM + stop", + ] +) + +STT_SVG = _svg_pipeline( + "Fluxo STT", + ["WAV base", "Variacao audio", "VAD proxy", "Sofya STT", "Prefixo WER", "CSV Markdown"], +) + +TTS_SVG = _svg_pipeline( + "Fluxo TTS", + ["Texto teste", "xAI TTS", "PCM audio", "WAV gerado", "Metricas audio", "CSV Markdown"], +) + +E2E_SVG = _svg_pipeline( + "Fluxo E2E Bridge LiveKit", + ["Runner", "Bridge", "LiveKit", "Agent", "Sofya STT", "xAI TTS", "Bridge"], +) + + +def _json_default(value: Any) -> Any: + if isinstance(value, Path): + return str(value) + if is_dataclass(value): + return asdict(value) + return str(value) + + +def write_json(path: Path | str, payload: Any) -> Path: + out = Path(path) + out.parent.mkdir(parents=True, exist_ok=True) + out.write_text(json.dumps(payload, ensure_ascii=False, indent=2, default=_json_default), encoding="utf-8") + return out + + +def write_csv(path: Path | str, rows: list[dict[str, Any]]) -> Path: + out = Path(path) + out.parent.mkdir(parents=True, exist_ok=True) + fieldnames: list[str] = [] + for row in rows: + for key in row: + if key not in fieldnames: + fieldnames.append(key) + with out.open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=fieldnames) + writer.writeheader() + for row in rows: + writer.writerow({key: row.get(key, "") for key in fieldnames}) + return out + + +def write_mermaid_files(report_dir: Path | str) -> dict[str, str]: + base_dir = Path(report_dir) / "diagrams" + base_dir.mkdir(parents=True, exist_ok=True) + files = { + "stt_mermaid": base_dir / "stt_flow.mmd", + "tts_mermaid": base_dir / "tts_flow.mmd", + "e2e_mermaid": base_dir / "e2e_flow.mmd", + "stt_svg": base_dir / "stt_flow.svg", + "tts_svg": base_dir / "tts_flow.svg", + "e2e_svg": base_dir / "e2e_flow.svg", + } + files["stt_mermaid"].write_text(STT_MERMAID + "\n", encoding="utf-8") + files["tts_mermaid"].write_text(TTS_MERMAID + "\n", encoding="utf-8") + files["e2e_mermaid"].write_text(E2E_MERMAID + "\n", encoding="utf-8") + files["stt_svg"].write_text(STT_SVG, encoding="utf-8") + files["tts_svg"].write_text(TTS_SVG, encoding="utf-8") + files["e2e_svg"].write_text(E2E_SVG, encoding="utf-8") + return {key: value.as_posix() for key, value in files.items()} + + +def _status(value: bool) -> str: + return "APROVADO" if value else "FALHOU" + + +def _human_bool(value: Any) -> str: + if value is True: + return "sim" + if value is False: + return "nao" + return str(value) + + +def _translate_error(value: Any) -> str: + text = str(value or "") + translations = { + "missing expected ready/stop/audio/timeline signal": ( + "faltou algum sinal esperado: ready, stop, audio ou timeline" + ), + } + return translations.get(text, text) + + +def _md_table(rows: list[dict[str, Any]], columns: list[str | tuple[str, str]]) -> str: + if not rows: + return "_Nenhuma linha._" + keys = [item[0] if isinstance(item, tuple) else item for item in columns] + labels = [item[1] if isinstance(item, tuple) else item for item in columns] + header = "| " + " | ".join(labels) + " |" + sep = "| " + " | ".join("---" for _ in labels) + " |" + body = [] + for row in rows: + values = [] + for key in keys: + value = row.get(key, "") + if key == "description" and not value: + value = scenario_description(str(row.get("scenario") or "")) + if key == "passed": + value = _status(value is True) + elif isinstance(value, bool): + value = _human_bool(value) + elif key == "error": + value = _translate_error(value) + text = str(value).replace("\n", " ").replace("|", "\\|") + values.append(text) + body.append("| " + " | ".join(values) + " |") + return "\n".join([header, sep, *body]) + + +def render_markdown_report(summary: dict[str, Any]) -> str: + stt_rows = summary.get("stt_results", []) + tts_rows = summary.get("tts_results", []) + e2e_rows = summary.get("e2e_results", []) + startup_rows = list((summary.get("startup_checks") or {}).values()) + synthetic = bool(summary.get("baseline", {}).get("synthetic")) + overall_passed = bool(summary.get("passed")) + + lines = [ + "# Relatorio do Teste Local de Estresse", + "", + "## Resumo Executivo", + "", + f"- Status geral: **{_status(overall_passed)}**", + f"- Inicio: `{summary.get('started_at', '')}`", + f"- Duracao: `{summary.get('duration_ms', 0)} ms`", + f"- Audio base: `{summary.get('baseline', {}).get('path', '')}`", + f"- Texto esperado: `{summary.get('expected_text', '')}`", + f"- Texto sintetizado: `{summary.get('synthesis_text') or summary.get('expected_text', '')}`", + f"- Chamadas STT: `{summary.get('stt_total_calls', len(stt_rows))}`", + f"- Prefixo monitorado: `{summary.get('stt_prefix_text', '')}`", + ] + vad_proxy = summary.get("vad_proxy") or {} + if vad_proxy: + lines.append( + "- VAD proxy: " + f"`threshold={vad_proxy.get('threshold_dbfs')} dBFS`, " + f"`prefix_padding={vad_proxy.get('prefix_padding_ms')} ms`, " + f"`min_speech={vad_proxy.get('min_speech_ms')} ms`" + ) + diagnostics = summary.get("diagnostics") or {} + if diagnostics: + lines.append( + "- Diagnostico audio/VAD: " + f"`STT_DUMP_DIR={diagnostics.get('stt_dump_dir', '')}`, " + f"`FLOW_LOG_VAD_DECISIONS={diagnostics.get('flow_log_vad_decisions', '')}`, " + f"`FLOW_LOG_VAD_ACTIVITY={diagnostics.get('flow_log_vad_activity', '')}`" + ) + if synthetic: + lines.extend( + [ + "", + "> Aviso: esta execucao usou um audio base sintetico gerado pelo xAI TTS. Quando houver uma gravacao humana, substitua com `STRESS_AUDIO`.", + ] + ) + if startup_rows: + lines.extend( + [ + "", + "## Prontidao Local", + "", + _md_table( + startup_rows, + [ + ("service", "servico"), + ("url", "url"), + ("status", "status"), + ("attempts", "tentativas"), + ("detail", "detalhe"), + ], + ), + ] + ) + + lines.extend( + [ + "", + "## Cenarios STT", + "", + _md_table( + stt_rows, + [ + ("scenario", "cenario"), + ("description", "descricao"), + ("passed", "status"), + ("prefix_ok", "prefixo_ok"), + ("vad_proxy_low_start_risk", "risco_inicio"), + ("vad_proxy_first_voice_ms", "1a_voz_ms"), + ("vad_proxy_unrecovered_prefix_ms", "prefixo_descoberto_ms"), + ("wer", "WER"), + ("duration_ms", "duracao_ms"), + ("rms_dbfs", "rms_dbfs"), + ("transcript", "transcricao"), + ("missing_terms", "termos_ausentes"), + ("error", "erro"), + ], + ), + "", + "## Cenarios TTS", + "", + _md_table( + tts_rows, + [ + ("scenario", "cenario"), + ("description", "descricao"), + ("passed", "status"), + ("duration_ms", "duracao_ms"), + ("latency_ms", "latencia_ms"), + ("rms_dbfs", "rms_dbfs"), + ("clipping_ratio", "taxa_clipping"), + ("audio_path", "audio"), + ("error", "erro"), + ], + ), + "", + "## Cenarios E2E Bridge/LiveKit", + "", + _md_table( + e2e_rows, + [ + ("scenario", "cenario"), + ("description", "descricao"), + ("passed", "status"), + ("ready_received", "recebeu_ready"), + ("stop_received", "recebeu_stop"), + ("non_silent_frames", "frames_com_audio"), + ("timeline_transcript", "transcricao_timeline"), + ("audio_path", "audio"), + ("error", "erro"), + ], + ), + "", + "## Fluxo STT", + "", + "Imagem estatica para preview no VSCode:", + "", + "![Fluxo STT](diagrams/stt_flow.svg)", + "", + "Versao textual:", + "", + "1. O runner carrega o WAV base.", + "2. Gera uma variacao de audio para o cenario.", + "3. Calcula um VAD proxy por energia RMS para estimar risco de inicio fraco.", + "4. Envia o audio para o Sofya STT HTTP.", + "5. Normaliza a transcricao retornada.", + "6. Verifica prefixo esperado, WER e termos criticos.", + "7. Escreve a linha correspondente no CSV e no Markdown.", + "", + f"Arquivo Mermaid separado: `{summary.get('artifacts', {}).get('stt_mermaid', '')}`", + "", + "Codigo Mermaid:", + "", + "```mermaid", + STT_MERMAID, + "```", + "", + "## Fluxo TTS", + "", + "Imagem estatica para preview no VSCode:", + "", + "![Fluxo TTS](diagrams/tts_flow.svg)", + "", + "Versao textual:", + "", + "1. O runner envia um texto de teste para o xAI TTS.", + "2. O xAI retorna audio PCM pelo websocket.", + "3. O runner salva o audio como WAV.", + "4. Calcula duracao, RMS e clipping.", + "5. Escreve a linha correspondente no CSV e no Markdown.", + "", + f"Arquivo Mermaid separado: `{summary.get('artifacts', {}).get('tts_mermaid', '')}`", + "", + "Codigo Mermaid:", + "", + "```mermaid", + TTS_MERMAID, + "```", + "", + "## Fluxo E2E", + "", + "Imagem estatica para preview no VSCode:", + "", + "![Fluxo E2E](diagrams/e2e_flow.svg)", + "", + "Versao textual:", + "", + "1. O runner abre websocket no Bridge e envia `start`.", + "2. O Bridge publica o audio do cliente no LiveKit.", + "3. O Agent recebe o audio e chama Sofya STT.", + "4. O backend fake gera a resposta de negocio.", + "5. O Agent chama xAI TTS para sintetizar a resposta.", + "6. O audio volta pelo LiveKit para o Bridge.", + "7. O Bridge devolve audio PCM e, ao final ideal, envia `stop`.", + "", + f"Arquivo Mermaid separado: `{summary.get('artifacts', {}).get('e2e_mermaid', '')}`", + "", + "Codigo Mermaid:", + "", + "```mermaid", + E2E_MERMAID, + "```", + "", + "## Artefatos", + "", + f"- Summary JSON: `{summary.get('artifacts', {}).get('summary_json', '')}`", + f"- CSV STT: `{summary.get('artifacts', {}).get('stt_csv', '')}`", + f"- CSV TTS: `{summary.get('artifacts', {}).get('tts_csv', '')}`", + f"- CSV E2E: `{summary.get('artifacts', {}).get('e2e_csv', '')}`", + f"- Recorte da timeline: `{summary.get('artifacts', {}).get('timeline_excerpt', '')}`", + f"- Mermaid STT: `{summary.get('artifacts', {}).get('stt_mermaid', '')}`", + f"- Mermaid TTS: `{summary.get('artifacts', {}).get('tts_mermaid', '')}`", + f"- Mermaid E2E: `{summary.get('artifacts', {}).get('e2e_mermaid', '')}`", + f"- SVG STT: `{summary.get('artifacts', {}).get('stt_svg', '')}`", + f"- SVG TTS: `{summary.get('artifacts', {}).get('tts_svg', '')}`", + f"- SVG E2E: `{summary.get('artifacts', {}).get('e2e_svg', '')}`", + ] + ) + return "\n".join(lines) + "\n" + + +def write_markdown_report(path: Path | str, summary: dict[str, Any]) -> Path: + out = Path(path) + out.parent.mkdir(parents=True, exist_ok=True) + out.write_text(render_markdown_report(summary), encoding="utf-8") + return out diff --git a/src/app/tools/local_stresstest/runner.py b/src/app/tools/local_stresstest/runner.py new file mode 100644 index 0000000..114ee18 --- /dev/null +++ b/src/app/tools/local_stresstest/runner.py @@ -0,0 +1,755 @@ +from __future__ import annotations + +import asyncio +import json +import os +import time +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Awaitable, Callable +from urllib.parse import urlsplit, urlunsplit + +import aiohttp +import httpx +import websockets +from dotenv import load_dotenv +from livekit.agents.types import DEFAULT_API_CONNECT_OPTIONS + +from app.livekit.adapters.xai_tts import OraclexAITTS +from app.providers.stt_config import override_config +from app.providers.stt_internal_livekit import InternalHTTPSTT, InternalSTTConfig, extract_text_from_transcript +from app.tools.local_stresstest.audio import ( + AudioSample, + audio_metrics, + build_variations, + pcm_to_wav_bytes, + read_wav_mono16, + vad_proxy_metrics, + write_wav, +) +from app.tools.local_stresstest.report import ( + render_markdown_report, + write_csv, + write_json, + write_markdown_report, + write_mermaid_files, +) +from app.tools.local_stresstest.scenarios import scenario_description +from app.tools.local_stresstest.text import compare_text, normalize_text +from app.tools.local_stresstest.timeline import excerpt_timeline, timeline_has_error, wait_for_timeline + + +DEFAULT_EXPECTED_TEXT = ( + "Por que os serviços Aia, Hexa e Banca aparecem com valor na Danf " + "se eles são incluídos no meu plano" +) +DEFAULT_CRITICAL_TERMS = ("por", "que", "aia", "hexa", "banca", "danf", "incluídos", "plano") +DEFAULT_SHORT_TTS_TEXT = "Teste de voz do xAI para validar áudio sintetizado." +DEFAULT_LONG_TTS_TEXT = ( + "Este é um teste operacional do TTS xAI para validar latência, duração, volume e entrega " + "de áudio antes do relatório local." +) +TERMINAL_EVENTS = { + "bridge_failed", + "stop_stt_unavailable", + "stop_tts_unavailable", + "stop_bridge_failed", + "stop_agent_runtime_unavailable", +} + + +@dataclass(frozen=True, slots=True) +class StressConfig: + env_file: Path + report_dir: Path + expected_text: str + synthesis_text: str + critical_terms: tuple[str, ...] + stt_wer_threshold: float + bridge_url: str + bridge_health_url: str + agent_health_url: str + startup_wait_s: float + startup_poll_s: float + skip_local_wait: bool + repeat: int + concurrency: int + e2e_turns: int + e2e_timeout_s: float + stress_audio: Path | None + prefix_text: str + prefix_words: int + vad_proxy_threshold_dbfs: float + vad_proxy_prefix_padding_ms: int + vad_proxy_min_speech_ms: int + + +def _env_str(name: str, default: str = "") -> str: + return str(os.getenv(name, default) or default).strip() + + +def _env_int(name: str, default: int) -> int: + try: + return int(_env_str(name, str(default))) + except ValueError: + return default + + +def _env_float(name: str, default: float) -> float: + try: + return float(_env_str(name, str(default))) + except ValueError: + return default + + +def _env_bool(name: str, default: bool = False) -> bool: + value = _env_str(name) + if not value: + return default + return value.lower() in {"1", "true", "yes", "on", "sim"} + + +def _first_words(text: str, count: int) -> str: + words = [item for item in str(text or "").strip().split() if item] + return " ".join(words[: max(1, count)]) + + +def _derive_bridge_health_url(bridge_url: str) -> str: + explicit = _env_str("STRESS_BRIDGE_HEALTH_URL") + if explicit: + return explicit + parsed = urlsplit(bridge_url) + if not parsed.netloc: + return "http://127.0.0.1:8000/health" + scheme = "https" if parsed.scheme == "wss" else "http" + return urlunsplit((scheme, parsed.netloc, "/health", "", "")) + + +def _derive_agent_health_url() -> str: + return _env_str( + "STRESS_AGENT_HEALTH_URL", + f"http://127.0.0.1:{_env_int('AGENT_SERVER_PORT', 18081)}/", + ) + + +def load_config() -> StressConfig: + env_file = Path(_env_str("STRESS_ENV_FILE", ".env.dev")) + if env_file.exists(): + load_dotenv(env_file, override=False) + + report_dir = Path(_env_str("STRESS_REPORT_DIR", ".run/local-stresstest")) + os.environ.setdefault("STT_DUMP_DIR", str(report_dir / "stt_dumps")) + os.environ.setdefault("FLOW_LOG_VAD_DECISIONS", "1") + os.environ.setdefault("FLOW_LOG_VAD_ACTIVITY", "1") + os.environ.setdefault("FLOW_LOG_VAD_ACTIVITY_MIN_PROB", "0.03") + audio_value = _env_str("STRESS_AUDIO") + critical_terms = tuple( + item.strip() + for item in _env_str("STRESS_CRITICAL_TERMS", ",".join(DEFAULT_CRITICAL_TERMS)).split(",") + if item.strip() + ) + bridge_url = _env_str("STRESS_BRIDGE_URL", "ws://127.0.0.1:8000/ws/agent") + expected_text = _env_str("STRESS_EXPECTED_TEXT", DEFAULT_EXPECTED_TEXT) + synthesis_text = _env_str("STRESS_SYNTH_TEXT", expected_text) + prefix_words = max(1, _env_int("STRESS_PREFIX_WORDS", 2)) + vad_prefix_padding_default_ms = round(_env_float("VAD_PREFIX_PADDING_DURATION", 1.5) * 1000) + return StressConfig( + env_file=env_file, + report_dir=report_dir, + expected_text=expected_text, + synthesis_text=synthesis_text, + critical_terms=critical_terms, + stt_wer_threshold=_env_float("STRESS_WER_THRESHOLD", 0.20), + bridge_url=bridge_url, + bridge_health_url=_derive_bridge_health_url(bridge_url), + agent_health_url=_derive_agent_health_url(), + startup_wait_s=max(0.0, _env_float("STRESS_STARTUP_WAIT_S", 180.0)), + startup_poll_s=max(0.5, _env_float("STRESS_STARTUP_POLL_S", 3.0)), + skip_local_wait=_env_bool("STRESS_SKIP_LOCAL_WAIT", False), + repeat=max(1, _env_int("STRESS_REPEAT", 1)), + concurrency=max(1, _env_int("STRESS_CONCURRENCY", 2)), + e2e_turns=max(1, _env_int("STRESS_E2E_TURNS", 3)), + e2e_timeout_s=max(5.0, _env_float("STRESS_E2E_TIMEOUT_S", 45.0)), + stress_audio=Path(audio_value) if audio_value else None, + prefix_text=_env_str("STRESS_PREFIX_TEXT", _first_words(expected_text, prefix_words)), + prefix_words=prefix_words, + vad_proxy_threshold_dbfs=_env_float("STRESS_VAD_PROXY_DBFS", -45.0), + vad_proxy_prefix_padding_ms=max( + 0, + _env_int("STRESS_VAD_PREFIX_PADDING_MS", vad_prefix_padding_default_ms), + ), + vad_proxy_min_speech_ms=max(0, _env_int("STRESS_VAD_MIN_SPEECH_MS", 100)), + ) + + +def _require_env(*names: str) -> str: + for name in names: + value = _env_str(name) + if value: + return value + raise RuntimeError(f"Missing required environment variable: {' or '.join(names)}") + + +ServiceProbe = Callable[[str], Awaitable[tuple[bool, str]]] +SleepFn = Callable[[float], Awaitable[None]] + + +async def _http_ready_probe(url: str) -> tuple[bool, str]: + timeout = httpx.Timeout(2.5, connect=1.0) + async with httpx.AsyncClient(timeout=timeout) as client: + response = await client.get(url) + ok = 200 <= response.status_code < 300 + return ok, f"HTTP {response.status_code}" + + +async def wait_for_local_services( + config: StressConfig, + *, + probe: ServiceProbe | None = None, + sleep: SleepFn = asyncio.sleep, +) -> dict[str, dict[str, Any]]: + services = { + "bridge": config.bridge_health_url, + "agent_runtime": config.agent_health_url, + } + states: dict[str, dict[str, Any]] = { + name: { + "service": name, + "url": url, + "ok": False, + "status": "pending", + "attempts": 0, + "detail": "", + } + for name, url in services.items() + } + if config.skip_local_wait: + for state in states.values(): + state.update({"ok": True, "status": "skipped", "detail": "STRESS_SKIP_LOCAL_WAIT habilitado"}) + return states + + active_probe = probe or _http_ready_probe + deadline = time.monotonic() + config.startup_wait_s + print( + "[local-stresstest] Aguardando servicos locais " + f"(timeout={config.startup_wait_s:.0f}s, poll={config.startup_poll_s:.1f}s)" + ) + while True: + for name, url in services.items(): + state = states[name] + if state["ok"]: + continue + state["attempts"] += 1 + try: + ok, detail = await active_probe(url) + except Exception as exc: + ok = False + detail = f"{type(exc).__name__}: {exc}" + state.update( + { + "ok": ok, + "status": "ready" if ok else "waiting", + "detail": detail, + } + ) + + if all(state["ok"] for state in states.values()): + print("[local-stresstest] Servicos locais prontos: bridge, agent_runtime") + return states + + if time.monotonic() >= deadline: + failed = "; ".join( + f"{name} em {state['url']} -> {state['detail'] or state['status']}" + for name, state in states.items() + if not state["ok"] + ) + raise RuntimeError( + f"Timeout aguardando servicos locais por {config.startup_wait_s:.0f}s: {failed}" + ) + + remaining = max(0.1, deadline - time.monotonic()) + await sleep(min(config.startup_poll_s, remaining)) + + +def _stt_config() -> InternalSTTConfig: + return InternalSTTConfig( + url=_require_env("STT_URL"), + api_key=_env_str("STT_KEY") or _env_str("STT_API_KEY", "unknow"), + language=_env_str("STT_LANG", "portuguese"), + timeout_s=_env_float("STRESS_STT_TIMEOUT_S", _env_float("STT_TIMEOUT_S", 30.0)), + min_prob_single_word=_env_float("STT_MIN_PROB_SINGLE_WORD", 0.10), + initial_prompt=_env_str("STT_INITIAL_PROMPT"), + config_override=_env_str("STT_CONFIG_OVERRIDE"), + output_mode=_env_str("STT_OUTPUT_MODE", "threshold_text"), + ) + + +def _stt_preprocessor_profile(cfg: InternalSTTConfig) -> str: + raw = (os.getenv("STT_FORCE_OVERRIDE_CONFIG", "") or cfg.config_override or "").strip() + payload: dict[str, Any] + if raw: + try: + parsed = json.loads(raw) + payload = parsed if isinstance(parsed, dict) else {} + except Exception: + return "custom_override_unparseable" + else: + payload = override_config + + preprocessors = payload.get("pre_processors") + if not isinstance(preprocessors, list): + return "none" + names = [ + str(item.get("strategy") or "").strip() + for item in preprocessors + if isinstance(item, dict) and str(item.get("strategy") or "").strip() + ] + return ",".join(names) if names else "none" + + +def _prefix_ok(*, expected_prefix: str, transcript: str) -> bool: + normalized_expected = normalize_text(expected_prefix) + normalized_actual = normalize_text(transcript) + return bool(normalized_expected) and normalized_actual.startswith(normalized_expected) + + +async def _collect_tts_bytes(tts: OraclexAITTS, text: str) -> bytes: + frame = await tts.synthesize(text, conn_options=DEFAULT_API_CONNECT_OPTIONS).collect() + return bytes(getattr(frame, "data", b"")) + + +async def _synthesize_xai_to_sample(text: str) -> AudioSample: + async with aiohttp.ClientSession() as session: + tts = OraclexAITTS( + api_key=_require_env("XAI_API_KEY"), + voice=_env_str("XAI_TTS_VOICE"), + language=_env_str("XAI_TTS_LANGUAGE", "pt-BR"), + http_session=session, + ) + try: + pcm = await _collect_tts_bytes(tts, text) + sample = AudioSample(name="xai_tts", pcm=pcm, sample_rate=tts.sample_rate, channels=tts.num_channels) + tmp_path = Path(".run/local-stresstest/.tmp_xai_collect.wav") + write_wav(tmp_path, sample) + return read_wav_mono16(tmp_path) + finally: + await tts.aclose() + + +async def prepare_baseline(config: StressConfig) -> dict[str, Any]: + assets_dir = config.report_dir / "assets" + assets_dir.mkdir(parents=True, exist_ok=True) + if config.stress_audio is not None: + sample = read_wav_mono16(config.stress_audio) + path = write_wav(assets_dir / "baseline_from_stress_audio.wav", sample) + return {"sample": sample, "path": path, "synthetic": False} + + sample = await _synthesize_xai_to_sample(config.synthesis_text) + path = write_wav(assets_dir / "synthetic_baseline.wav", sample) + return {"sample": sample, "path": path, "synthetic": True} + + +async def run_stt_scenario( + *, + config: StressConfig, + scenario: AudioSample, + output_path: Path, +) -> dict[str, Any]: + started = time.perf_counter() + metrics = audio_metrics(scenario.pcm, sample_rate=scenario.sample_rate, channels=scenario.channels) + vad_metrics = vad_proxy_metrics( + scenario, + threshold_dbfs=config.vad_proxy_threshold_dbfs, + prefix_padding_ms=config.vad_proxy_prefix_padding_ms, + min_speech_ms=config.vad_proxy_min_speech_ms, + ) + write_wav(output_path, scenario) + stt_cfg = _stt_config() + row: dict[str, Any] = { + "scenario": scenario.name, + "description": scenario_description(scenario.name), + "audio_path": str(output_path), + "duration_ms": metrics.duration_ms, + "rms_dbfs": metrics.rms_dbfs, + "peak_dbfs": metrics.peak_dbfs, + "clipping_ratio": metrics.clipping_ratio, + "vad_proxy_threshold_dbfs": vad_metrics.threshold_dbfs, + "vad_proxy_prefix_padding_ms": vad_metrics.prefix_padding_ms, + "vad_proxy_first_voice_ms": vad_metrics.first_voice_ms, + "vad_proxy_padded_start_ms": vad_metrics.padded_start_ms, + "vad_proxy_unrecovered_prefix_ms": vad_metrics.unrecovered_prefix_ms, + "vad_proxy_speech_ms": vad_metrics.speech_ms, + "vad_proxy_initial_200ms_dbfs": vad_metrics.initial_200ms_dbfs, + "vad_proxy_initial_500ms_dbfs": vad_metrics.initial_500ms_dbfs, + "vad_proxy_low_start_risk": vad_metrics.low_start_risk, + "expected_prefix": config.prefix_text, + "stt_preprocessors": _stt_preprocessor_profile(stt_cfg), + } + client = httpx.AsyncClient() + try: + stt = InternalHTTPSTT(stt_cfg, client=client) + event = await stt._post_internal_http( + scenario.pcm, + pcm_to_wav_bytes(scenario), + req_id=f"stress-{scenario.name}", + started_ns=time.time_ns(), + language=_env_str("STT_LANG", "portuguese"), + audio_duration_ms=metrics.duration_ms, + level_dbfs=metrics.rms_dbfs, + ) + transcript = event.alternatives[0].text if event.alternatives else "" + transcript = extract_text_from_transcript(transcript) + prefix_ok = _prefix_ok(expected_prefix=config.prefix_text, transcript=transcript) + comparison = compare_text( + expected=config.expected_text, + actual=transcript, + critical_terms=config.critical_terms, + ) + passed = bool(transcript) and comparison.terms_ok and comparison.wer <= config.stt_wer_threshold + row.update( + { + "passed": passed, + "transcript": transcript, + "actual_prefix": _first_words(transcript, config.prefix_words), + "prefix_ok": prefix_ok, + "wer": round(comparison.wer, 4), + "substitutions": comparison.substitutions, + "deletions": comparison.deletions, + "insertions": comparison.insertions, + "missing_terms": ",".join(comparison.missing_terms), + "latency_ms": round((time.perf_counter() - started) * 1000), + "error": "", + } + ) + except Exception as exc: + row.update( + { + "passed": False, + "transcript": "", + "actual_prefix": "", + "prefix_ok": False, + "wer": "", + "missing_terms": "", + "latency_ms": round((time.perf_counter() - started) * 1000), + "error": str(exc), + } + ) + finally: + await client.aclose() + return row + + +async def run_stt_suite(config: StressConfig, baseline: AudioSample) -> list[dict[str, Any]]: + audio_dir = config.report_dir / "stt_audio" + variations = [] + for repeat_idx in range(1, config.repeat + 1): + for variation in build_variations(baseline): + if config.repeat > 1: + variations.append( + AudioSample( + name=f"{variation.name}_r{repeat_idx}", + pcm=variation.pcm, + sample_rate=variation.sample_rate, + channels=variation.channels, + ) + ) + else: + variations.append(variation) + rows: list[dict[str, Any]] = [] + semaphore = asyncio.Semaphore(config.concurrency) + + async def _run(item: AudioSample, idx: int) -> dict[str, Any]: + async with semaphore: + return await run_stt_scenario( + config=config, + scenario=item, + output_path=audio_dir / f"{idx:02d}_{item.name}.wav", + ) + + tasks = [ + _run(variation, idx) + for idx, variation in enumerate(variations, start=1) + ] + for result in await asyncio.gather(*tasks): + rows.append(result) + return rows + + +async def run_tts_suite(config: StressConfig) -> list[dict[str, Any]]: + scenarios = [ + ("short", DEFAULT_SHORT_TTS_TEXT), + ("long", DEFAULT_LONG_TTS_TEXT), + ("reuse_1", DEFAULT_SHORT_TTS_TEXT), + ("reuse_2", DEFAULT_SHORT_TTS_TEXT), + ] + rows: list[dict[str, Any]] = [] + out_dir = config.report_dir / "tts_audio" + async with aiohttp.ClientSession() as session: + tts = OraclexAITTS( + api_key=_require_env("XAI_API_KEY"), + voice=_env_str("XAI_TTS_VOICE"), + language=_env_str("XAI_TTS_LANGUAGE", "pt-BR"), + http_session=session, + ) + try: + for name, text in scenarios: + started = time.perf_counter() + row: dict[str, Any] = {"scenario": name, "text_len": len(text)} + row["description"] = scenario_description(name) + try: + pcm = await _collect_tts_bytes(tts, text) + sample = AudioSample(name=name, pcm=pcm, sample_rate=tts.sample_rate, channels=tts.num_channels) + path = write_wav(out_dir / f"{name}.wav", sample) + normalized = read_wav_mono16(path) + metrics = audio_metrics(normalized.pcm) + expected_min_ms = max(250, len(text) * 18) + expected_max_ms = max(2500, len(text) * 180) + passed = ( + bool(pcm) + and metrics.audible + and metrics.clipping_ratio <= _env_float("STRESS_TTS_MAX_CLIPPING", 0.02) + and expected_min_ms <= metrics.duration_ms <= expected_max_ms + ) + row.update( + { + "passed": passed, + "audio_path": str(path), + "duration_ms": metrics.duration_ms, + "latency_ms": round((time.perf_counter() - started) * 1000), + "rms_dbfs": metrics.rms_dbfs, + "peak_dbfs": metrics.peak_dbfs, + "clipping_ratio": metrics.clipping_ratio, + "error": "", + } + ) + except Exception as exc: + row.update( + { + "passed": False, + "audio_path": "", + "duration_ms": 0, + "latency_ms": round((time.perf_counter() - started) * 1000), + "rms_dbfs": "", + "clipping_ratio": "", + "error": str(exc), + } + ) + rows.append(row) + finally: + await tts.aclose() + return rows + + +def _start_payload(config: StressConfig) -> dict[str, Any]: + stamp = datetime.now(timezone.utc).strftime("%Y%m%d%H%M%S") + return { + "type": "start", + "data": { + "agent": "conta", + "ani": "5511999990000", + "gsm": "5511999990000", + "routerCallKeyDay": datetime.now().strftime("%Y%m%d"), + "routerCallKey": f"RCK-stress-{stamp}", + "callIdGed": f"GED-stress-{stamp}", + "session_id": f"stress-{stamp}", + "agentData": {"idFatura": "FAT-STRESS-001"}, + }, + "audioFormat": {"encoding": "linear16", "sampleRateHz": 16000, "channels": 1}, + "callConfig": { + "agentBackend": "remote_ws_fake", + "stt": {"provider": "internal_http"}, + "tts": {"provider": "xai"}, + }, + } + + +async def _send_pcm_turns(ws: Any, sample: AudioSample, *, turns: int) -> None: + frame_bytes = round(sample.sample_rate * 20 / 1000) * 2 + silence = b"\x00" * frame_bytes + for _ in range(turns): + for offset in range(0, len(sample.pcm), frame_bytes): + chunk = sample.pcm[offset : offset + frame_bytes] + if len(chunk) < frame_bytes: + chunk = chunk + b"\x00" * (frame_bytes - len(chunk)) + await ws.send(chunk) + await asyncio.sleep(0.02) + for _ in range(60): + await ws.send(silence) + await asyncio.sleep(0.02) + + +async def run_e2e_suite(config: StressConfig, baseline: AudioSample) -> list[dict[str, Any]]: + scenario = AudioSample(name="clean", pcm=baseline.pcm) + started = time.perf_counter() + row: dict[str, Any] = { + "scenario": scenario.name, + "description": scenario_description(scenario.name), + "ready_received": False, + "stop_received": False, + "non_silent_frames": 0, + "timeline_transcript": "", + "audio_path": "", + "error": "", + } + received_audio = bytearray() + text_messages: list[dict[str, Any]] = [] + room = "" + + try: + async with websockets.connect(config.bridge_url, max_size=2**24) as ws: + await ws.send(json.dumps(_start_payload(config), ensure_ascii=False)) + first = await asyncio.wait_for(ws.recv(), timeout=10.0) + if isinstance(first, bytes): + raise RuntimeError("Expected ready/stop text message, got binary") + payload = json.loads(first) + text_messages.append(payload) + if payload.get("type") == "stop": + row["error"] = json.dumps(payload, ensure_ascii=False) + row["stop_received"] = True + return [dict(row, passed=False)] + if payload.get("type") != "ready": + raise RuntimeError(f"Expected ready, got {payload}") + row["ready_received"] = True + room = str(payload.get("room") or "") + + sender = asyncio.create_task(_send_pcm_turns(ws, scenario, turns=config.e2e_turns)) + deadline = time.monotonic() + config.e2e_timeout_s + while time.monotonic() < deadline: + try: + message = await asyncio.wait_for(ws.recv(), timeout=0.5) + except asyncio.TimeoutError: + if sender.done(): + continue + continue + if isinstance(message, bytes): + received_audio.extend(message) + if audio_metrics(message).audible: + row["non_silent_frames"] += 1 + continue + + payload = json.loads(message) + text_messages.append(payload) + if payload.get("type") == "stop": + row["stop_received"] = True + break + + sender.cancel() + await asyncio.gather(sender, return_exceptions=True) + + if received_audio: + audio_path = config.report_dir / "e2e_audio" / f"{scenario.name}_bridge_response.wav" + write_wav(audio_path, AudioSample(name="bridge_response", pcm=bytes(received_audio))) + row["audio_path"] = str(audio_path) + + timeline_dir = Path(_env_str("CALL_TIMELINE_DIR", "./timeline")) + records = wait_for_timeline(timeline_dir / f"{room}.jsonl", timeout_s=5.0) if room else [] + transcripts = [ + str(record.get("text") or record.get("transcript") or "") + for record in records + if str(record.get("event") or "") in {"user_transcript_final", "stt_final"} + ] + row["timeline_transcript"] = " | ".join(item for item in transcripts if item) + row["timeline_path"] = str(timeline_dir / f"{room}.jsonl") if room else "" + row["timeline_error"] = timeline_has_error(records) + row["text_messages"] = json.dumps(text_messages, ensure_ascii=False) + terminal_error = any( + str(message.get("data", {}).get("status") or "") in TERMINAL_EVENTS + for message in text_messages + if isinstance(message, dict) + ) + row["passed"] = bool( + row["ready_received"] + and row["stop_received"] + and row["non_silent_frames"] > 0 + and row["timeline_transcript"] + and not row["timeline_error"] + and not terminal_error + ) + if not row["passed"] and not row["error"]: + row["error"] = "missing expected ready/stop/audio/timeline signal" + + excerpt_path = config.report_dir / "timeline_excerpt.jsonl" + excerpt_path.parent.mkdir(parents=True, exist_ok=True) + with excerpt_path.open("w", encoding="utf-8") as handle: + for record in excerpt_timeline(records): + handle.write(json.dumps(record, ensure_ascii=False)) + handle.write("\n") + except Exception as exc: + row["passed"] = False + row["error"] = str(exc) + + row["latency_ms"] = round((time.perf_counter() - started) * 1000) + return [row] + + +def _artifact_paths(config: StressConfig) -> dict[str, str]: + artifacts = { + "summary_json": str(config.report_dir / "summary.json"), + "stt_csv": str(config.report_dir / "stt_results.csv"), + "tts_csv": str(config.report_dir / "tts_results.csv"), + "e2e_csv": str(config.report_dir / "e2e_results.csv"), + "timeline_excerpt": str(config.report_dir / "timeline_excerpt.jsonl"), + "report_md": str(config.report_dir / "report.md"), + } + artifacts.update(write_mermaid_files(config.report_dir)) + return artifacts + + +async def run() -> int: + config = load_config() + config.report_dir.mkdir(parents=True, exist_ok=True) + started = time.perf_counter() + started_at = datetime.now(timezone.utc).isoformat() + + startup_checks = await wait_for_local_services(config) + baseline_info = await prepare_baseline(config) + baseline = baseline_info["sample"] + stt_results = await run_stt_suite(config, baseline) + tts_results = await run_tts_suite(config) + e2e_results = await run_e2e_suite(config, baseline) + + artifacts = _artifact_paths(config) + write_csv(artifacts["stt_csv"], stt_results) + write_csv(artifacts["tts_csv"], tts_results) + write_csv(artifacts["e2e_csv"], e2e_results) + + summary = { + "started_at": started_at, + "duration_ms": round((time.perf_counter() - started) * 1000), + "expected_text": config.expected_text, + "synthesis_text": config.synthesis_text, + "critical_terms": list(config.critical_terms), + "wer_threshold": config.stt_wer_threshold, + "stt_scenarios_per_repeat": len(build_variations(baseline)), + "stt_total_calls": len(stt_results), + "stt_prefix_text": config.prefix_text, + "vad_proxy": { + "threshold_dbfs": config.vad_proxy_threshold_dbfs, + "prefix_padding_ms": config.vad_proxy_prefix_padding_ms, + "min_speech_ms": config.vad_proxy_min_speech_ms, + "note": "Proxy por energia RMS; nao substitui o Silero VAD do LiveKit.", + }, + "diagnostics": { + "stt_dump_dir": _env_str("STT_DUMP_DIR"), + "flow_log_vad_decisions": _env_str("FLOW_LOG_VAD_DECISIONS"), + "flow_log_vad_activity": _env_str("FLOW_LOG_VAD_ACTIVITY"), + "flow_log_vad_activity_min_prob": _env_str("FLOW_LOG_VAD_ACTIVITY_MIN_PROB"), + }, + "startup_checks": startup_checks, + "baseline": { + "path": str(baseline_info["path"]), + "synthetic": bool(baseline_info["synthetic"]), + "metrics": audio_metrics(baseline.pcm), + }, + "stt_results": stt_results, + "tts_results": tts_results, + "e2e_results": e2e_results, + "artifacts": artifacts, + } + summary["passed"] = all(row.get("passed") is True for row in [*stt_results, *tts_results, *e2e_results]) + write_json(artifacts["summary_json"], summary) + write_markdown_report(artifacts["report_md"], summary) + + print(render_markdown_report(summary)) + return 0 if summary["passed"] else 1 diff --git a/src/app/tools/local_stresstest/scenarios.py b/src/app/tools/local_stresstest/scenarios.py new file mode 100644 index 0000000..05d3a8b --- /dev/null +++ b/src/app/tools/local_stresstest/scenarios.py @@ -0,0 +1,47 @@ +from __future__ import annotations + + +SCENARIO_DESCRIPTIONS = { + "clean": "Audio base sem degradacao; serve como linha de referencia do STT.", + "low_volume": "Mesmo audio com ganho reduzido para simular fala baixa ou captura distante.", + "very_low_volume": "Mesmo audio com ganho muito baixo para testar limite de audibilidade do STT.", + "high_volume": "Mesmo audio com ganho elevado, sem clipping esperado, para simular fala forte.", + "clipped_high_volume": "Audio com ganho alto o suficiente para provocar clipping e distorcao.", + "leading_trailing_silence": "Audio com silencio antes e depois da fala para validar recorte/VAD.", + "short_leading_silence": "Pequeno silencio antes da fala para validar inicio com pouca margem.", + "long_leading_silence": "Silencio longo antes da fala para observar trim e tempo de ativacao.", + "noise_snr_20": "Audio com ruido leve, SNR 20 dB, para medir robustez inicial.", + "noise_snr_15": "Audio com ruido moderado, SNR 15 dB, para observar degradacao.", + "noise_snr_10": "Audio com ruido forte, SNR 10 dB, para identificar limite de tolerancia.", + "pre_noise_300ms": "Ruido curto antes da fala para simular ambiente aberto antes do enunciado.", + "pre_noise_800ms": "Ruido mais longo antes da fala para testar se preprocessors confundem inicio.", + "telephony_profile": "Audio passado por perfil telefonico simples via u-law para simular compressao.", + "telephony_low_volume": "Perfil telefonico combinado com volume baixo.", + "initial_fade_in_250ms": "Ataque inicial suavizado por 250 ms para simular primeira palavra fraca.", + "initial_fade_in_500ms": "Ataque inicial suavizado por 500 ms para estressar captura do prefixo.", + "initial_fade_in_900ms": "Ataque inicial suavizado por 900 ms, perto do prefix padding local.", + "initial_dip_300ms": "Primeiros 300 ms com ganho baixo, depois volume normal.", + "initial_dip_700ms": "Primeiros 700 ms com ganho baixo, simulando pergunta iniciada fraca.", + "prefix_300ms_15pct": "Prefixo de 300 ms muito baixo para observar corte de inicio curto.", + "prefix_600ms_20pct": "Prefixo de 600 ms baixo para simular 'Por que' sem energia.", + "prefix_900ms_25pct": "Prefixo de 900 ms baixo, limite interessante para padding de 1s.", + "low_prefix_noise_600ms": "Prefixo baixo com ruido, combinando fala fraca e ambiente ruidoso.", + "low_prefix_telephony_700ms": "Prefixo baixo com perfil telefonico para simular URA/telefonia.", + "short": "Sintese curta no xAI TTS para validar resposta basica, volume e latencia.", + "long": "Sintese longa no xAI TTS para validar estabilidade e duracao plausivel.", + "reuse_1": "Primeira repeticao curta para observar reuso da conexao websocket do xAI.", + "reuse_2": "Segunda repeticao curta para confirmar reuso continuo da conexao websocket do xAI.", +} + + +def scenario_base_name(name: str) -> str: + value = str(name or "").strip() + if "_r" in value: + prefix, suffix = value.rsplit("_r", 1) + if suffix.isdigit(): + return prefix + return value + + +def scenario_description(name: str) -> str: + return SCENARIO_DESCRIPTIONS.get(scenario_base_name(name), "") diff --git a/src/app/tools/local_stresstest/text.py b/src/app/tools/local_stresstest/text.py new file mode 100644 index 0000000..f5d0ac7 --- /dev/null +++ b/src/app/tools/local_stresstest/text.py @@ -0,0 +1,106 @@ +from __future__ import annotations + +import re +import unicodedata +from dataclasses import dataclass + + +_TOKEN_RE = re.compile(r"[^0-9a-zA-Z]+") + + +@dataclass(frozen=True, slots=True) +class TextComparison: + expected: str + actual: str + normalized_expected: str + normalized_actual: str + wer: float + substitutions: int + deletions: int + insertions: int + missing_terms: tuple[str, ...] + + @property + def terms_ok(self) -> bool: + return not self.missing_terms + + +def normalize_text(text: str) -> str: + value = unicodedata.normalize("NFD", str(text or "").strip().lower()) + value = "".join(ch for ch in value if unicodedata.category(ch) != "Mn") + value = _TOKEN_RE.sub(" ", value) + return " ".join(value.split()) + + +def _edit_counts(expected_tokens: list[str], actual_tokens: list[str]) -> tuple[int, int, int]: + rows = len(expected_tokens) + 1 + cols = len(actual_tokens) + 1 + dp: list[list[tuple[int, int, int, int]]] = [ + [(0, 0, 0, 0) for _ in range(cols)] + for _ in range(rows) + ] + + for i in range(1, rows): + cost, sub, dele, ins = dp[i - 1][0] + dp[i][0] = (cost + 1, sub, dele + 1, ins) + for j in range(1, cols): + cost, sub, dele, ins = dp[0][j - 1] + dp[0][j] = (cost + 1, sub, dele, ins + 1) + + for i in range(1, rows): + for j in range(1, cols): + if expected_tokens[i - 1] == actual_tokens[j - 1]: + candidates = [dp[i - 1][j - 1]] + else: + cost, sub, dele, ins = dp[i - 1][j - 1] + candidates = [(cost + 1, sub + 1, dele, ins)] + + cost, sub, dele, ins = dp[i - 1][j] + candidates.append((cost + 1, sub, dele + 1, ins)) + + cost, sub, dele, ins = dp[i][j - 1] + candidates.append((cost + 1, sub, dele, ins + 1)) + + dp[i][j] = min(candidates, key=lambda item: (item[0], item[1], item[2], item[3])) + + _, substitutions, deletions, insertions = dp[-1][-1] + return substitutions, deletions, insertions + + +def word_error_rate(expected: str, actual: str) -> tuple[float, int, int, int]: + expected_tokens = normalize_text(expected).split() + actual_tokens = normalize_text(actual).split() + if not expected_tokens: + return (0.0 if not actual_tokens else 1.0, 0, 0, len(actual_tokens)) + + substitutions, deletions, insertions = _edit_counts(expected_tokens, actual_tokens) + wer = (substitutions + deletions + insertions) / len(expected_tokens) + return wer, substitutions, deletions, insertions + + +def compare_text( + *, + expected: str, + actual: str, + critical_terms: list[str] | tuple[str, ...], +) -> TextComparison: + normalized_expected = normalize_text(expected) + normalized_actual = normalize_text(actual) + actual_tokens = set(normalized_actual.split()) + missing_terms = tuple( + term + for term in (normalize_text(item) for item in critical_terms) + if term and term not in actual_tokens + ) + wer, substitutions, deletions, insertions = word_error_rate(expected, actual) + return TextComparison( + expected=expected, + actual=actual, + normalized_expected=normalized_expected, + normalized_actual=normalized_actual, + wer=wer, + substitutions=substitutions, + deletions=deletions, + insertions=insertions, + missing_terms=missing_terms, + ) diff --git a/src/app/tools/local_stresstest/timeline.py b/src/app/tools/local_stresstest/timeline.py new file mode 100644 index 0000000..39ed34f --- /dev/null +++ b/src/app/tools/local_stresstest/timeline.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +import json +import time +from pathlib import Path +from typing import Any + + +IMPORTANT_EVENTS = { + "ready_sent", + "client_audio_first_frame_received", + "client_audio_first_frame_published", + "client_audio_activity_detected", + "agent_join", + "user_transcript_final", + "stt_final", + "stt_http_completed", + "tts_stage_started", + "tts_stage_result", + "done_packet_received", + "stop_sent", + "bridge_failed", + "call_end", +} + +ERROR_EVENTS = { + "bridge_failed", + "stt_http_failed", + "stt_http_failed_nonfatal", + "tts_error", + "tts_stage_failed", + "agent_reconnect_exhausted", +} + + +def read_timeline(path: Path | str) -> list[dict[str, Any]]: + timeline_path = Path(path) + if not timeline_path.exists(): + return [] + records: list[dict[str, Any]] = [] + for line in timeline_path.read_text(encoding="utf-8").splitlines(): + if not line.strip(): + continue + try: + payload = json.loads(line) + except json.JSONDecodeError: + continue + if isinstance(payload, dict): + records.append(payload) + return records + + +def wait_for_timeline(path: Path | str, *, timeout_s: float = 5.0) -> list[dict[str, Any]]: + deadline = time.monotonic() + timeout_s + timeline_path = Path(path) + while time.monotonic() < deadline: + records = read_timeline(timeline_path) + if records: + return records + time.sleep(0.1) + return read_timeline(timeline_path) + + +def excerpt_timeline(records: list[dict[str, Any]]) -> list[dict[str, Any]]: + return [ + record + for record in records + if str(record.get("event") or "") in IMPORTANT_EVENTS + ] + + +def timeline_has_error(records: list[dict[str, Any]]) -> bool: + return any(str(record.get("event") or "") in ERROR_EVENTS for record in records) diff --git a/src/app/tools/oci_audio_download.py b/src/app/tools/oci_audio_download.py new file mode 100644 index 0000000..8201348 --- /dev/null +++ b/src/app/tools/oci_audio_download.py @@ -0,0 +1,229 @@ +from __future__ import annotations + +import argparse +import re +import sys +import zipfile +from dataclasses import dataclass, replace +from datetime import date +from pathlib import Path +from typing import Any, Iterable + +from app.utils.stt_audio_upload import OCIUploadConfig, _oci_client, _oci_upload_config_from_env + + +_SAFE_SESSION_ID = re.compile(r"^[A-Za-z0-9_.-]+$") + + +@dataclass(frozen=True, slots=True) +class DownloadResult: + prefix: str + output_dir: Path + object_count: int + downloaded_count: int + cached_count: int + total_bytes: int + zip_path: Path | None + + +def _parse_date(value: str) -> str: + try: + return date.fromisoformat(value).isoformat() + except ValueError as exc: + raise argparse.ArgumentTypeError("use uma data no formato AAAA-MM-DD") from exc + + +def _parse_session_id(value: str) -> str: + session_id = value.strip() + if not session_id or not _SAFE_SESSION_ID.fullmatch(session_id): + raise argparse.ArgumentTypeError( + "session-id deve conter apenas letras, numeros, ponto, hifen ou underscore" + ) + return session_id + + +def segments_prefix(day: str, session_id: str) -> str: + return f"{day}/{session_id}/" + + +def entire_calls_prefix(day: str) -> str: + return f"{day}/entire_call/" + + +def _list_objects(client: Any, config: OCIUploadConfig, prefix: str) -> list[Any]: + objects: list[Any] = [] + start: str | None = None + while True: + response = client.list_objects( + namespace_name=config.namespace, + bucket_name=config.bucket, + prefix=prefix, + start=start, + fields="name,size,etag", + ) + objects.extend( + item + for item in response.data.objects + if getattr(item, "name", "") and not item.name.endswith("/") + ) + start = getattr(response.data, "next_start_with", None) + if not start: + return objects + + +def _safe_relative_name(object_name: str, prefix: str) -> Path: + if not object_name.startswith(prefix): + raise ValueError(f"objeto fora do prefixo solicitado: {object_name}") + relative = Path(object_name[len(prefix) :]) + if not relative.parts or relative.is_absolute() or ".." in relative.parts: + raise ValueError(f"nome de objeto inseguro: {object_name}") + return relative + + +def _response_chunks(response: Any) -> Iterable[bytes]: + raw = getattr(getattr(response, "data", None), "raw", None) + if raw is not None and hasattr(raw, "stream"): + yield from raw.stream(1024 * 1024, decode_content=False) + return + data = getattr(response, "data", None) + content = getattr(data, "content", data) + if isinstance(content, bytes): + yield content + return + raise TypeError("resposta do Object Storage nao contem um corpo legivel") + + +def _create_zip(output_dir: Path, zip_path: Path) -> None: + zip_path.parent.mkdir(parents=True, exist_ok=True) + temporary = zip_path.with_suffix(zip_path.suffix + ".part") + try: + with zipfile.ZipFile(temporary, "w", compression=zipfile.ZIP_DEFLATED) as archive: + for path in sorted(output_dir.rglob("*")): + if path.is_file() and not path.name.endswith(".part"): + archive.write(path, path.relative_to(output_dir)) + temporary.replace(zip_path) + finally: + temporary.unlink(missing_ok=True) + + +def download_prefix( + *, + client: Any, + config: OCIUploadConfig, + prefix: str, + output_dir: Path, + zip_path: Path | None, +) -> DownloadResult: + objects = _list_objects(client, config, prefix) + if not objects: + raise FileNotFoundError( + f"nenhum objeto encontrado em bucket={config.bucket} prefix={prefix}" + ) + + output_dir.mkdir(parents=True, exist_ok=True) + downloaded = cached = total_bytes = 0 + print(f"bucket={config.bucket} prefix={prefix} objetos={len(objects)}", flush=True) + + for index, item in enumerate(objects, 1): + destination = output_dir / _safe_relative_name(item.name, prefix) + destination.parent.mkdir(parents=True, exist_ok=True) + expected_size = int(getattr(item, "size", 0) or 0) + if destination.exists() and expected_size and destination.stat().st_size == expected_size: + cached += 1 + total_bytes += expected_size + print(f"[{index}/{len(objects)}] cache {item.name}", flush=True) + continue + + temporary = destination.with_suffix(destination.suffix + ".part") + try: + response = client.get_object( + namespace_name=config.namespace, + bucket_name=config.bucket, + object_name=item.name, + ) + with temporary.open("wb") as handle: + for chunk in _response_chunks(response): + handle.write(chunk) + actual_size = temporary.stat().st_size + if expected_size and actual_size != expected_size: + raise IOError( + f"tamanho divergente para {item.name}: " + f"esperado={expected_size} recebido={actual_size}" + ) + temporary.replace(destination) + finally: + temporary.unlink(missing_ok=True) + + downloaded += 1 + total_bytes += destination.stat().st_size + print(f"[{index}/{len(objects)}] baixado {item.name}", flush=True) + + if zip_path is not None: + _create_zip(output_dir, zip_path) + print(f"zip={zip_path} bytes={zip_path.stat().st_size}", flush=True) + + return DownloadResult( + prefix, output_dir, len(objects), downloaded, cached, total_bytes, zip_path + ) + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description="Baixa audios do OCI Object Storage usados pelo TIA." + ) + parser.add_argument("--bucket", help="sobrescreve BUCKET_NAME (dev, fqa ou prd)") + parser.add_argument( + "--output-root", + type=Path, + default=Path("recordings/object_storage"), + help="diretorio raiz (padrao: recordings/object_storage)", + ) + parser.add_argument("--no-zip", action="store_true", help="nao gera o ZIP final") + commands = parser.add_subparsers(dest="command", required=True) + segments = commands.add_parser("segments", help="segmentos de uma chamada") + segments.add_argument("--date", required=True, type=_parse_date) + segments.add_argument("--session-id", required=True, type=_parse_session_id) + entire = commands.add_parser("entire-calls", help="chamadas completas de um dia") + entire.add_argument("--date", required=True, type=_parse_date) + return parser + + +def main(argv: list[str] | None = None) -> int: + args = _parser().parse_args(argv) + config = _oci_upload_config_from_env() + if config is None: + print("erro: configuracao OCI BUCKET_* incompleta", file=sys.stderr) + return 2 + if args.bucket: + config = replace(config, bucket=args.bucket.strip()) + + if args.command == "segments": + prefix = segments_prefix(args.date, args.session_id) + output_dir = args.output_root / args.date / "segments" / args.session_id + zip_path = None if args.no_zip else args.output_root / f"{args.date}_{args.session_id}_segments.zip" + else: + prefix = entire_calls_prefix(args.date) + output_dir = args.output_root / args.date / "entire_call" + zip_path = None if args.no_zip else args.output_root / f"{args.date}_entire_calls.zip" + + try: + result = download_prefix( + client=_oci_client(config), + config=config, + prefix=prefix, + output_dir=output_dir, + zip_path=zip_path, + ) + except Exception as exc: + print(f"erro: {exc}", file=sys.stderr) + return 1 + print( + f"concluido objetos={result.object_count} baixados={result.downloaded_count} " + f"cache={result.cached_count} bytes={result.total_bytes}", + flush=True, + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/src/app/tools/oci_logs_cli_search.py b/src/app/tools/oci_logs_cli_search.py new file mode 100644 index 0000000..75772a5 --- /dev/null +++ b/src/app/tools/oci_logs_cli_search.py @@ -0,0 +1,678 @@ +from __future__ import annotations + +import argparse +import configparser +import csv +import json +import re +import shutil +import subprocess +import sys +import tempfile +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from zoneinfo import ZoneInfo + + +DEFAULT_LOG_SCOPE = ( + "ocid1.compartment.oc1..aaaaaaaau5raupdqclblygmdikwcqt2mglgrfkgdmca6kleven5maaqau22q/" + "ocid1.loggroup.oc1.sa-saopaulo-1.amaaaaaaaehl73aaude3eznx7pbzbilwwzay5ou2mcccr5aj5tz6rlbyamua" +) +DEFAULT_REGION = "sa-saopaulo-1" +TIA_AGENT_SUBJECT = "tim-ai-atend-agnt-integ-tia-app-agent/0.log" +TIA_BRIDGE_SUBJECT = "tim-ai-atend-agnt-integ-tia-app-bridge/0.log" +CONTAS_NAMESPACE = "agnt-ai-atendimento-contas" + +VAD_STT_PATTERNS = [ + "step=vad_activity", + "step=vad_speech_start", + "step=vad_speech_end", + "step=vad_user_pause", + "step=vad_interrupt_check", + "step=stt_audio_upload", + "step=stt_final", + "USER_STATE", + "interrupt", + "tts_start", + "tts_done", + "BOT_SAY", + "agent_input_snapshot", +] + +IMPORTANT_EVENTS = { + "vad_speech_start", + "vad_speech_end", + "vad_user_pause", + "vad_interrupt_check", + "stt_audio_upload_enqueued", + "stt_audio_upload_done", + "stt_final", + "USER_STATE", + "tts_start", + "tts_done", + "BOT_SAY_HANDLE", + "BOT_SAY_RESULT", + "agent_input_snapshot", +} + +LOG_LINE_RE = re.compile(r"^(?P\S+)\s+(?Pstdout|stderr)\s+(?P[FP])\s(?P.*)$") +KV_RE = re.compile(r"(?P[a-zA-Z_][a-zA-Z0-9_]*)=(?P[^|]+)") + + +def clean(value: str | None) -> str: + return (value or "").strip() + + +def quote(value: str) -> str: + return value.replace("\\", "\\\\").replace("'", "\\'") + + +def contains_wildcard(field: str, value: str) -> str: + return f"{field} = '*{quote(value)}*'" + + +def equals_value(field: str, value: str) -> str: + return f"{field} = '{quote(value)}'" + + +def starts_with_value(field: str, value: str) -> str: + return f"{field} = '{quote(value)}*'" + + +def contains_any(field: str, values: list[str]) -> str: + unique_values: list[str] = [] + for value in values: + value = clean(value) + if value and value not in unique_values: + unique_values.append(value) + return "(" + " or ".join(contains_wildcard(field, value) for value in unique_values) + ")" + + +def _filter_attr(item: Any, name: str, default: str = "") -> str: + if isinstance(item, dict): + return clean(str(item.get(name, default) or "")) + return clean(str(getattr(item, name, default) or "")) + + +def _simple_filter_values(field_name: str, value: str) -> tuple[str, list[str]]: + query_field = "data.message" + if field_name in {"subject", "container"}: + query_field = "subject" + + values = [value] + if field_name == "event" and value and "=" not in value: + values.append(f"step={value}") + if field_name in {"session_id", "message_id"}: + compact = hyphenless(value) + if compact: + values.append(compact) + return query_field, values + + +def simple_filter_expression(item: Any) -> str: + field_name = _filter_attr(item, "field", "message") + operator = _filter_attr(item, "operator", "contains") + value = _filter_attr(item, "value") + if not value: + return "" + + query_field, values = _simple_filter_values(field_name, value) + if operator == "not_contains": + return "not " + contains_any(query_field, values) + if operator == "equals": + return "(" + " or ".join(equals_value(query_field, item_value) for item_value in values) + ")" + if operator == "starts_with": + return "(" + " or ".join(starts_with_value(query_field, item_value) for item_value in values) + ")" + return contains_any(query_field, values) + + +def simple_filters_group(filters: list[Any]) -> str: + parts: list[str] = [] + for item in filters: + expression = simple_filter_expression(item) + if not expression: + continue + connector = _filter_attr(item, "connector", "and").lower() + connector = "or" if connector == "or" else "and" + if not parts: + parts.append(expression) + else: + parts.append(f"{connector} {expression}") + if not parts: + return "" + return "(" + " ".join(parts) + ")" + + +def hyphenless(value: str) -> str | None: + clean_value = clean(value) + if re.match(r"^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$", clean_value): + return clean_value.replace("-", "") + return None + + +def normalize_time(value: str, time_zone: str) -> str: + raw = clean(value) + if not raw: + raise ValueError("Horario vazio.") + raw = raw.replace(" ", "T") + if raw.endswith("Z"): + parsed = datetime.fromisoformat(raw[:-1] + "+00:00") + else: + parsed = datetime.fromisoformat(raw) + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=ZoneInfo(time_zone)) + return parsed.astimezone(timezone.utc).isoformat().replace("+00:00", "Z") + + +def windows_path_to_wsl(value: str) -> str: + match = re.match(r"^([a-zA-Z]):[\\/](.*)$", value.strip()) + if not match: + return value + drive = match.group(1).lower() + rest = match.group(2).replace("\\", "/") + return f"/mnt/{drive}/{rest}" + + +def pem_fingerprint(path: Path) -> str | None: + try: + from cryptography.hazmat.primitives import serialization + from cryptography.hazmat.primitives.serialization import load_pem_private_key + + key = load_pem_private_key(path.read_bytes(), password=None) + der = key.public_key().public_bytes( + serialization.Encoding.DER, + serialization.PublicFormat.SubjectPublicKeyInfo, + ) + return ":".join(f"{byte:02x}" for byte in __import__("hashlib").md5(der).digest()) + except Exception: + return None + + +def find_matching_key_file(config_file: Path, key_file: str, fingerprint: str) -> str | None: + expected = clean(fingerprint).lower() + if not expected: + return None + + current = Path(windows_path_to_wsl(key_file)).expanduser() + if current.exists() and pem_fingerprint(current) == expected: + return str(current) + + search_dirs = [] + for candidate_dir in (current.parent, config_file.parent): + if candidate_dir and candidate_dir.exists() and candidate_dir not in search_dirs: + search_dirs.append(candidate_dir) + + matches: list[Path] = [] + for directory in search_dirs: + for candidate in directory.glob("*.pem"): + if pem_fingerprint(candidate) == expected: + matches.append(candidate) + + unique_matches = sorted({match.resolve(): match for match in matches}.values()) + if len(unique_matches) == 1: + return str(unique_matches[0]) + return None + + +def config_profile_value(config: configparser.RawConfigParser, profile: str | None, key: str) -> str: + profile_name = profile or "DEFAULT" + if profile_name == "DEFAULT": + return config.defaults().get(key, "") + return config.get(profile_name, key, fallback="") + + +def prepare_cli_config(config_file: str | None, profile: str | None) -> str | None: + if not config_file or sys.platform.startswith("win"): + return config_file + + source = Path(config_file).expanduser() + if not source.exists(): + return config_file + + original = source.read_text(encoding="utf-8") + + def replace_path(match: re.Match[str]) -> str: + key = match.group("key") + value = match.group("value").strip() + converted = windows_path_to_wsl(value) + return f"{key}{converted}" + + converted = re.sub( + r"^(?P\s*(?:key_file|security_token_file)\s*=\s*)(?P[a-zA-Z]:[^\r\n]+)$", + replace_path, + original, + flags=re.MULTILINE, + ) + parser = configparser.RawConfigParser() + parser.read_string(converted) + key_file = config_profile_value(parser, profile, "key_file") + fingerprint = config_profile_value(parser, profile, "fingerprint") + matching_key_file = find_matching_key_file(source, key_file, fingerprint) + if matching_key_file and matching_key_file != key_file: + converted = re.sub( + r"^(?P\s*key_file\s*=\s*)(?P[^\r\n]+)$", + lambda match: f"{match.group('key')}{matching_key_file}", + converted, + count=1, + flags=re.MULTILINE, + ) + if converted == original: + return config_file + + handle = tempfile.NamedTemporaryFile("w", encoding="utf-8", prefix="oci_config_wsl_", delete=False) + with handle: + handle.write(converted) + return handle.name + + +def build_search_query(args: argparse.Namespace) -> str: + subjects = [] + app_scope = getattr(args, "app_scope", "tia") + if app_scope in {"tia", "all"}: + if args.include_agent: + subjects.append(contains_wildcard("subject", TIA_AGENT_SUBJECT)) + if args.include_bridge: + subjects.append(contains_wildcard("subject", TIA_BRIDGE_SUBJECT)) + if app_scope in {"contas", "all"}: + subjects.append(contains_wildcard("logContent", CONTAS_NAMESPACE)) + if not subjects: + raise ValueError("Selecione pelo menos uma aplicacao ou componente para consultar.") + + where_parts = ["(" + " or ".join(subjects) + ")"] + if not args.show_health: + where_parts.append("not (" + contains_wildcard("data.message", "GET /health") + ")") + + if args.event_preset == "vad-stt": + where_parts.append(contains_any("data.message", VAD_STT_PATTERNS)) + + id_filter_parts: list[str] = [] + for arg_name, field_name in ( + ("session_id", "data.message"), + ("message_id", "data.message"), + ): + value = clean(getattr(args, arg_name)) + if not value: + continue + values = [value] + compact = hyphenless(value) + if compact: + values.append(compact) + id_filter_parts.append(contains_any(field_name, values)) + + for arg_name in ("room", "protocol", "call_id_ged"): + value = clean(getattr(args, arg_name)) + if value: + id_filter_parts.append(contains_wildcard("data.message", value)) + + if id_filter_parts: + id_filter = "(" + " and ".join(id_filter_parts) + ")" + if args.event_preset == "vad-stt" and not args.strict_id_filter: + raw_vad_filter = contains_any( + "data.message", + [ + "step=vad_activity", + "step=vad_speech_start", + "step=vad_speech_end", + "step=vad_user_pause", + "step=vad_interrupt_check", + "USER_STATE", + ], + ) + where_parts.append(f"({id_filter} or {raw_vad_filter})") + else: + where_parts.append(id_filter) + + for text in args.text or []: + if clean(text): + where_parts.append(contains_wildcard("data.message", text)) + + simple_group = simple_filters_group(list(getattr(args, "simple_filters", []) or [])) + if simple_group: + where_parts.append(simple_group) + + lines = [f'search "{clean(args.log_scope)}"'] + lines.extend(f"| where {part}" for part in where_parts) + lines.append("| select datetime, subject, data.message as message") + lines.append(f"| sort by datetime {args.sort}") + return "\n".join(lines) + + +def next_page_token(payload: dict[str, Any]) -> str | None: + for key in ("opc-next-page", "opcNextPage", "nextPage"): + value = payload.get(key) + if value: + return str(value) + headers = payload.get("headers") + if isinstance(headers, dict): + for key in ("opc-next-page", "opcNextPage", "nextPage"): + value = headers.get(key) + if value: + return str(value) + data = payload.get("data") + if isinstance(data, dict): + for key in ("opc-next-page", "opcNextPage", "nextPage"): + value = data.get(key) + if value: + return str(value) + return None + + +def extract_results(payload: dict[str, Any]) -> list[Any]: + data = payload.get("data", payload) + if isinstance(data, dict): + results = data.get("results") + if isinstance(results, list): + return results + if isinstance(data, list): + return data + return [] + + +def extract_summary(payload: dict[str, Any]) -> dict[str, Any]: + data = payload.get("data", payload) + if isinstance(data, dict) and isinstance(data.get("summary"), dict): + return data["summary"] + return {} + + +def run_oci_page(args: argparse.Namespace, query: str, time_start: str, time_end: str, page: str | None) -> dict[str, Any]: + python_dir = Path(sys.executable).parent + venv_dir = Path(sys.prefix) / ("Scripts" if sys.platform.startswith("win") else "bin") + oci_candidates = [ + shutil.which("oci"), + shutil.which("oci.exe"), + str(python_dir / "oci"), + str(python_dir / "oci.exe"), + str(python_dir / "oci.cmd"), + str(venv_dir / "oci"), + str(venv_dir / "oci.exe"), + str(venv_dir / "oci.cmd"), + ] + oci_bin = args.oci_bin or next((candidate for candidate in oci_candidates if candidate and Path(candidate).exists()), None) + if not oci_bin: + raise RuntimeError( + "Nao encontrei o binario 'oci'. Instale a OCI CLI, instale 'oci-cli' neste venv, " + "ou passe --oci-bin com o caminho do executavel." + ) + + command = [ + oci_bin, + "logging-search", + "search-logs", + "--region", + args.region, + "--time-start", + time_start, + "--time-end", + time_end, + "--search-query", + query, + "--limit", + str(args.limit), + "--output", + "json", + ] + if args.profile: + command.extend(["--profile", args.profile]) + if args.config_file: + command.extend(["--config-file", args.config_file]) + if page: + command.extend(["--page", page]) + + env = { + **dict(__import__("os").environ), + "SUPPRESS_LABEL_WARNING": "True", + "OCI_CLI_SUPPRESS_FILE_PERMISSIONS_WARNING": "True", + } + completed = subprocess.run(command, check=False, capture_output=True, text=True, encoding="utf-8", env=env) + if completed.returncode != 0: + raise RuntimeError( + "OCI CLI falhou com exit code " + f"{completed.returncode}.\nSTDERR:\n{completed.stderr.strip()}\nSTDOUT:\n{completed.stdout.strip()}" + ) + try: + return json.loads(completed.stdout) + except json.JSONDecodeError as exc: + raise RuntimeError(f"OCI CLI nao retornou JSON valido:\n{completed.stdout[:2000]}") from exc + + +def result_data(result: Any) -> dict[str, Any]: + if not isinstance(result, dict): + return {} + data = result.get("data") or result + if not isinstance(data, dict): + return {} + log_content = data.get("logContent") if isinstance(data.get("logContent"), dict) else {} + message = data.get("message") or data.get("data.message") or "" + if not message and isinstance(log_content.get("data"), dict): + message = (log_content.get("data") or {}).get("message", "") + return { + "datetime": data.get("datetime") or log_content.get("time") or "", + "subject": data.get("subject") or log_content.get("subject") or "", + "message": message, + } + + +def field(body: str, key: str) -> str | None: + match = re.search(rf"(?:^|\|)\s*{re.escape(key)}=([^|]+)", body) + if match: + return match.group(1).strip() + return None + + +def parse_message(message: str) -> dict[str, Any]: + parsed: dict[str, Any] = { + "log_time": "", + "body": message or "", + "event": "", + "fields": {}, + } + match = LOG_LINE_RE.match(message or "") + if match: + parsed.update(match.groupdict()) + body = str(parsed.get("body") or "") + fields = {item.group("key"): item.group("value").strip() for item in KV_RE.finditer(body)} + parsed["fields"] = fields + + event = fields.get("step") + if not event and body.startswith("{"): + try: + event = json.loads(body).get("tipo_evento") or "json" + except json.JSONDecodeError: + event = "json" + if not event: + event = body.split("|", 1)[0].strip() + parsed["event"] = event[:100] + return parsed + + +def normalize_row(result: Any) -> dict[str, Any]: + data = result_data(result) + parsed = parse_message(str(data.get("message") or "")) + return { + "datetime": str(data.get("datetime") or ""), + "subject": str(data.get("subject") or ""), + "log_time": parsed["log_time"], + "event": parsed["event"], + "body": parsed["body"], + "fields": parsed["fields"], + "message": data.get("message") or "", + } + + +def compact_row(row: dict[str, Any]) -> str: + fields = row.get("fields") or {} + details = [] + for key in ( + "old", + "new", + "message_id", + "duration_ms", + "original_audio_ms", + "sent_audio_ms", + "probability", + "speaking", + "speech_duration_ms", + "silence_ms", + "raw_speech_ms", + "raw_silence_ms", + "eligible", + "text", + "reason", + "stage", + ): + if key in fields: + details.append(f"{key}={fields[key]}") + if not details and row["body"].startswith("{"): + try: + payload = json.loads(row["body"]) + for key in ("tipo_evento", "message_id", "callid", "session_id", "mensagem"): + if payload.get(key): + details.append(f"{key}={payload[key]}") + except json.JSONDecodeError: + pass + if not details: + details.append(row["body"][:220]) + return " | ".join(details) + + +def print_summary(rows: list[dict[str, Any]]) -> None: + print(f"linhas={len(rows)}") + if rows: + print(f"primeira={rows[0].get('log_time') or rows[0].get('datetime')}") + print(f"ultima={rows[-1].get('log_time') or rows[-1].get('datetime')}") + + vad_rows = [row for row in rows if row.get("event") == "vad_activity"] + if vad_rows: + def probability(row: dict[str, Any]) -> float: + try: + return float((row.get("fields") or {}).get("probability") or 0) + except ValueError: + return 0.0 + + max_row = max(vad_rows, key=probability) + print("") + print("vad_activity") + print(f"count={len(vad_rows)}") + print(f"max_prob={probability(max_row):.3f} at={max_row.get('log_time') or max_row.get('datetime')}") + for threshold in (0.30, 0.20, 0.10): + print(f"above_{threshold:.2f}={sum(1 for row in vad_rows if probability(row) >= threshold)}") + print(f"speaking_true={sum(1 for row in vad_rows if (row.get('fields') or {}).get('speaking') == 'True')}") + + print("") + print("timeline") + for row in rows: + event = str(row.get("event") or "") + if event in IMPORTANT_EVENTS or "interrupt" in event: + print(f"{row.get('log_time') or row.get('datetime')} | {event} | {compact_row(row)}") + + +def write_json(path: Path, payload: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8") + + +def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=["datetime", "log_time", "event", "subject", "body"]) + writer.writeheader() + for row in rows: + writer.writerow({key: row.get(key, "") for key in writer.fieldnames or []}) + + +def parse_args(argv: list[str]) -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Busca logs da OCI via OCI CLI e resume eventos VAD/STT.") + parser.add_argument("--time-start", "--start", required=True, help="Inicio da janela. Use UTC ou informe --timezone.") + parser.add_argument("--time-end", "--end", required=True, help="Fim da janela. Use UTC ou informe --timezone.") + parser.add_argument("--timezone", default="UTC", help="Timezone usado quando os horarios vierem sem offset. Ex: America/Sao_Paulo.") + parser.add_argument("--region", default=DEFAULT_REGION) + parser.add_argument("--log-scope", default=DEFAULT_LOG_SCOPE) + parser.add_argument("--profile", default=None) + parser.add_argument("--config-file", default=None) + parser.add_argument("--oci-bin", default=None) + parser.add_argument("--limit", type=int, default=1000) + parser.add_argument("--max-pages", type=int, default=20) + parser.add_argument("--max-results", type=int, default=10000) + parser.add_argument("--sort", choices=["asc", "desc"], default="asc") + parser.add_argument("--include-agent", action=argparse.BooleanOptionalAction, default=True) + parser.add_argument("--include-bridge", action=argparse.BooleanOptionalAction, default=False) + parser.add_argument("--show-health", action="store_true") + parser.add_argument("--event-preset", choices=["vad-stt", "all"], default="vad-stt") + parser.add_argument( + "--strict-id-filter", + action="store_true", + help="Aplica session/call/message id de forma estrita. Por padrao, VAD cru tambem entra na janela.", + ) + parser.add_argument("--session-id") + parser.add_argument("--message-id") + parser.add_argument("--room") + parser.add_argument("--protocol") + parser.add_argument("--call-id-ged") + parser.add_argument("--text", action="append", help="Filtro livre em data.message. Pode repetir.") + parser.add_argument("--output", default=".tmp/oci_logs_cli_search.json") + parser.add_argument("--csv", dest="csv_path", default=None) + parser.add_argument("--dry-run", action="store_true") + parser.add_argument("--no-summary", action="store_true") + return parser.parse_args(argv) + + +def main(argv: list[str] | None = None) -> int: + args = parse_args(argv or sys.argv[1:]) + try: + time_start = normalize_time(args.time_start, args.timezone) + time_end = normalize_time(args.time_end, args.timezone) + query = build_search_query(args) + args.config_file = prepare_cli_config(args.config_file, args.profile) + except Exception as exc: + print(f"Erro nos argumentos: {exc}", file=sys.stderr) + return 2 + + if args.dry_run: + print("time_start=" + time_start) + print("time_end=" + time_end) + print(query) + return 0 + + all_results: list[Any] = [] + summaries: list[dict[str, Any]] = [] + page = None + try: + for page_number in range(1, args.max_pages + 1): + payload = run_oci_page(args, query, time_start, time_end, page) + results = extract_results(payload) + summaries.append(extract_summary(payload)) + all_results.extend(results) + print(f"pagina={page_number} resultados={len(results)} total={len(all_results)}") + page = next_page_token(payload) + if not page or len(all_results) >= args.max_results: + break + except Exception as exc: + print(str(exc), file=sys.stderr) + return 1 + + rows = [normalize_row(result) for result in all_results] + output_payload = { + "query": query, + "time_start": time_start, + "time_end": time_end, + "result_count": len(rows), + "summaries": summaries, + "rows": rows, + } + output_path = Path(args.output) + write_json(output_path, output_payload) + print(f"json={output_path}") + if args.csv_path: + csv_path = Path(args.csv_path) + write_csv(csv_path, rows) + print(f"csv={csv_path}") + if not args.no_summary: + print("") + print_summary(rows) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/src/app/tools/oci_logs_viewer.html b/src/app/tools/oci_logs_viewer.html new file mode 100644 index 0000000..90d4dc3 --- /dev/null +++ b/src/app/tools/oci_logs_viewer.html @@ -0,0 +1,1282 @@ + + + + + + OCI TIA Logs + + + +
+

OCI TIA Logs

+ Pronto +
+ + + +
+ + +
+
+
+ Nenhuma busca executada + +
+
+ + + Pagina 0/0 + + + + +
+
+ +
+ + + + + + + + + + +
TempoContainerEventoMensagem
+
+ +
+
Selecione uma linha para ver a mensagem completa.
+
+
+
+ + + + diff --git a/src/app/tools/oci_logs_viewer.py b/src/app/tools/oci_logs_viewer.py new file mode 100644 index 0000000..d5c2620 --- /dev/null +++ b/src/app/tools/oci_logs_viewer.py @@ -0,0 +1,508 @@ +from __future__ import annotations + +import argparse +import json +import re +import os +import shutil +import subprocess +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Literal + +from fastapi import FastAPI, HTTPException +from fastapi.responses import FileResponse +from pydantic import BaseModel, Field + +from app.tools.oci_logs_cli_search import ( + DEFAULT_LOG_SCOPE, + DEFAULT_REGION, + CONTAS_NAMESPACE, + TIA_AGENT_SUBJECT, + TIA_BRIDGE_SUBJECT, + build_search_query as build_cli_search_query, + extract_results, + extract_summary, + next_page_token, + normalize_row, + prepare_cli_config, + run_oci_page, + windows_path_to_wsl, +) + + +LOG_LINE_RE = re.compile(r"^(?P\S+)\s+(?Pstdout|stderr)\s+(?P[FP])\s(?P.*)$") +KV_RE = re.compile(r"(?P[a-zA-Z_][a-zA-Z0-9_]*)=(?P[^|]+)") +UUID_RE = re.compile(r"^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$") + + +class CliAuth(BaseModel): + mode: Literal["cli"] = "cli" + config_file: str = "~/.oci/config" + profile: str = "DEFAULT" + region: str = DEFAULT_REGION + oci_bin: str | None = None + + +class SimpleFilter(BaseModel): + connector: Literal["and", "or"] = "and" + field: Literal[ + "message", + "subject", + "container", + "event", + "session_id", + "message_id", + "call_id_ged", + "room", + "protocol", + ] = "message" + operator: Literal["contains", "not_contains", "equals", "starts_with"] = "contains" + value: str = "" + + +class QueryPreviewRequest(BaseModel): + log_scope: str = DEFAULT_LOG_SCOPE + app_scope: Literal["tia", "contas", "all"] = "tia" + sort: Literal["asc", "desc"] = "asc" + include_agent: bool = True + include_bridge: bool = True + hide_health: bool = True + event_preset: Literal["all", "vad-stt"] = "all" + strict_id_filter: bool = False + session_id: str | None = None + message_id: str | None = None + room: str | None = None + protocol: str | None = None + call_id_ged: str | None = None + text: str | None = None + simple_filters: list[SimpleFilter] = Field(default_factory=list) + + +class LogSearchRequest(QueryPreviewRequest): + auth: CliAuth = Field(default_factory=CliAuth) + time_start: str + time_end: str + limit: int = Field(default=500, ge=1, le=1000) + page: str | None = None + auto_page: bool = True + max_pages: int = Field(default=20, ge=1, le=50) + + +class CodexAnalyzeRequest(BaseModel): + prompt: str = Field(min_length=1, max_length=2_000_000) + codex_bin: str | None = None + model: str | None = None + timeout_seconds: int = Field(default=180, ge=30, le=900) + + +def _clean(value: str | None) -> str: + return (value or "").strip() + + +def _api_error(title: str, message: str, hint: str | None = None, raw: str | None = None) -> dict[str, str | None]: + return { + "title": title, + "message": message, + "hint": hint, + "raw": raw, + } + + +def _parse_datetime(value: str, field_name: str) -> datetime: + raw = value.strip() + if not raw: + raise HTTPException(status_code=400, detail=f"{field_name} vazio.") + if raw.endswith("Z"): + raw = raw[:-1] + "+00:00" + try: + parsed = datetime.fromisoformat(raw) + except ValueError as exc: + raise HTTPException(status_code=400, detail=f"{field_name} invalido: {value}") from exc + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=timezone.utc) + return parsed.astimezone(timezone.utc) + + +def _rfc3339(value: datetime) -> str: + return value.astimezone(timezone.utc).isoformat().replace("+00:00", "Z") + + +def _container_from_subject(subject: str) -> str: + match = re.search(r"/([^/]+)/0\.log$", subject or "") + return match.group(1) if match else "" + + +def _hyphenless(value: str) -> str | None: + clean = value.strip() + if UUID_RE.match(clean): + return clean.replace("-", "") + return None + + +def _query_args(filters: QueryPreviewRequest, auth: CliAuth | None = None, *, limit: int = 500) -> argparse.Namespace: + text = _clean(filters.text) + return argparse.Namespace( + log_scope=_clean(filters.log_scope) or DEFAULT_LOG_SCOPE, + app_scope=filters.app_scope, + sort=filters.sort, + include_agent=filters.include_agent, + include_bridge=filters.include_bridge, + show_health=not filters.hide_health, + event_preset=filters.event_preset, + strict_id_filter=filters.strict_id_filter, + session_id=_clean(filters.session_id) or None, + message_id=_clean(filters.message_id) or None, + room=_clean(filters.room) or None, + protocol=_clean(filters.protocol) or None, + call_id_ged=_clean(filters.call_id_ged) or None, + text=[text] if text else [], + simple_filters=[item.model_dump() for item in filters.simple_filters], + region=(auth.region if auth else DEFAULT_REGION), + profile=(auth.profile if auth else None), + config_file=(auth.config_file if auth else None), + oci_bin=(auth.oci_bin if auth else None), + limit=limit, + ) + + +def build_search_query(filters: QueryPreviewRequest) -> str: + return build_cli_search_query(_query_args(filters)) + + +def _ids_from_body(body: str) -> dict[str, str]: + ids: dict[str, str] = {} + for item in KV_RE.finditer(body): + key = item.group("key") + value = item.group("value").strip().strip('"').strip("'") + if key in { + "session_id", + "message_id", + "room", + "protocol", + "call_id_ged", + "phone_number", + "bridge", + "reason", + "ws_close_reason", + "ws_client_disconnected", + }: + ids[key] = value + + if body.startswith("{"): + try: + payload = json.loads(body) + except json.JSONDecodeError: + payload = {} + for key, payload_key in ( + ("session_id", "session_id"), + ("message_id", "message_id"), + ("call_id_ged", "callid"), + ): + value = payload.get(payload_key) + if value and key not in ids: + ids[key] = str(value) + + return ids + + +def _normalize_viewer_row(result: Any) -> dict[str, Any]: + row = normalize_row(result) + body = str(row.get("body") or "") + fields = row.get("fields") or {} + ids = _ids_from_body(body) + for key in ("session_id", "message_id", "room", "protocol", "call_id_ged", "phone_number", "reason"): + if fields.get(key) and key not in ids: + ids[key] = str(fields[key]) + return { + **row, + "container": _container_from_subject(str(row.get("subject") or "")), + "ids": ids, + "region": "", + } + + +def _summary_with_count(summary: dict[str, Any], count: int) -> dict[str, Any]: + return {**summary, "resultCount": count} + + +def _default_config_file() -> str: + home_config = Path("~/.oci/config").expanduser() + if home_config.exists(): + return str(home_config) + windows_configs = sorted(Path("/mnt/c/Users").glob("*/.oci/config")) if Path("/mnt/c/Users").exists() else [] + if len(windows_configs) == 1: + return str(windows_configs[0]) + return "~/.oci/config" + + +def _repo_root() -> Path: + return Path(__file__).resolve().parents[3] + + +def _candidate_codex_bins(explicit: str | None) -> list[str]: + candidates: list[str] = [] + if _clean(explicit): + raw = _clean(explicit) + candidates.append(raw) + if Path("/mnt/c").exists(): + candidates.append(windows_path_to_wsl(raw)) + + for name in ("codex", "codex.exe"): + found = shutil.which(name) + if found: + candidates.append(found) + + extension_roots = [Path.home() / ".vscode" / "extensions"] + windows_users = Path("/mnt/c/Users") + if windows_users.exists(): + extension_roots.extend(path / ".vscode" / "extensions" for path in windows_users.glob("*")) + + for root in extension_roots: + try: + candidates.extend( + str(path) + for path in sorted( + root.glob("openai.chatgpt-*/bin/windows-x86_64/codex.exe"), + reverse=True, + ) + ) + except OSError: + continue + + unique: list[str] = [] + seen: set[str] = set() + for candidate in candidates: + if candidate and candidate not in seen: + unique.append(candidate) + seen.add(candidate) + return unique + + +def _find_codex_bin(explicit: str | None) -> str | None: + for candidate in _candidate_codex_bins(explicit): + path = Path(candidate).expanduser() + if path.exists(): + return str(path) + found = shutil.which(candidate) + if found: + return found + return None + + +def _windows_path_from_wsl(value: str) -> str: + match = re.match(r"^/mnt/([a-zA-Z])/(.*)$", value) + if not match: + return value + drive = match.group(1).upper() + rest = match.group(2).replace("/", "\\") + return f"{drive}:\\{rest}" + + +def _codex_workdir_arg(codex_bin: str) -> str: + root = str(_repo_root()) + if codex_bin.lower().endswith(".exe"): + return _windows_path_from_wsl(root) + return root + + +def _codex_command(request: CodexAnalyzeRequest, codex_bin: str) -> list[str]: + command = [ + codex_bin, + "exec", + "--sandbox", + "read-only", + "--color", + "never", + "--skip-git-repo-check", + "--ephemeral", + "-C", + _codex_workdir_arg(codex_bin), + ] + model = _clean(request.model) + if model: + command.extend(["--model", model]) + command.append("-") + return command + + +app = FastAPI(title="OCI TIA Logs Viewer", version="2.0-cli") + + +@app.get("/") +def index() -> FileResponse: + return FileResponse(Path(__file__).with_name("oci_logs_viewer.html")) + + +@app.get("/api/defaults") +def defaults() -> dict[str, Any]: + return { + "log_scope": DEFAULT_LOG_SCOPE, + "region": DEFAULT_REGION, + "agent_subject": TIA_AGENT_SUBJECT, + "bridge_subject": TIA_BRIDGE_SUBJECT, + "contas_namespace": CONTAS_NAMESPACE, + "auth_mode": "cli", + "config_file": _default_config_file(), + } + + +@app.post("/api/query") +def preview_query(request: QueryPreviewRequest) -> dict[str, str]: + try: + query = build_search_query(request) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + return {"query": query} + + +@app.post("/api/search") +def search_logs(request: LogSearchRequest) -> dict[str, Any]: + time_start = _parse_datetime(request.time_start, "time_start") + time_end = _parse_datetime(request.time_end, "time_end") + if time_end <= time_start: + raise HTTPException(status_code=400, detail="time_end precisa ser maior que time_start.") + + query_request = QueryPreviewRequest( + **request.model_dump(exclude={"auth", "time_start", "time_end", "limit", "page", "auto_page", "max_pages"}) + ) + try: + query = build_search_query(query_request) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + cli_args = _query_args(query_request, request.auth, limit=request.limit) + try: + cli_args.config_file = prepare_cli_config(cli_args.config_file, cli_args.profile) + results: list[Any] = [] + summaries: list[dict[str, Any]] = [] + next_page = request.page + pages_fetched = 0 + while True: + payload = run_oci_page( + cli_args, + query, + _rfc3339(time_start), + _rfc3339(time_end), + next_page, + ) + results.extend(extract_results(payload)) + summaries.append(extract_summary(payload)) + pages_fetched += 1 + next_page = next_page_token(payload) + if request.page or not request.auto_page or not next_page or pages_fetched >= request.max_pages: + break + except RuntimeError as exc: + raise HTTPException( + status_code=502, + detail=_api_error( + "OCI CLI falhou", + str(exc), + "Confira o binario oci, profile/config, permissao IAM, regiao, escopo do log e sintaxe da query.", + ), + ) from exc + except Exception as exc: + raise HTTPException( + status_code=500, + detail=_api_error( + "Erro local ao consultar a OCI via CLI", + str(exc), + "Confira caminhos de config/PEM, WSL, rede/VPN e tente novamente.", + ), + ) from exc + + rows = [_normalize_viewer_row(item) for item in results] + summary = _summary_with_count(summaries[-1] if summaries else {}, len(rows)) + summary["pagesFetched"] = pages_fetched + summary["rowsFetched"] = len(rows) + return { + "query": query, + "time_start": _rfc3339(time_start), + "time_end": _rfc3339(time_end), + "rows": rows, + "summary": summary, + "headers": { + "opc_request_id": None, + "next_page": next_page, + }, + } + + +@app.post("/api/analyze") +def analyze_with_codex(request: CodexAnalyzeRequest) -> dict[str, Any]: + codex_bin = _find_codex_bin(request.codex_bin) + if not codex_bin: + raise HTTPException( + status_code=400, + detail=_api_error( + "Codex CLI nao encontrado", + "Nao encontrei o binario codex/codex.exe.", + "Preencha o caminho do Codex CLI na tela ou confira se ele esta no PATH.", + ), + ) + + command = _codex_command(request, codex_bin) + env = {**os.environ, "NO_COLOR": "1"} + try: + completed = subprocess.run( + command, + input=request.prompt, + text=True, + capture_output=True, + timeout=request.timeout_seconds, + cwd=_repo_root(), + env=env, + check=False, + ) + except subprocess.TimeoutExpired as exc: + raise HTTPException( + status_code=504, + detail=_api_error( + "Codex CLI excedeu o timeout", + f"A analise passou de {request.timeout_seconds}s.", + "Aumente o timeout ou reduza o numero de logs/char por log do pacote.", + ), + ) from exc + except OSError as exc: + raise HTTPException( + status_code=502, + detail=_api_error( + "Falha ao iniciar Codex CLI", + str(exc), + "Confira o caminho do binario Codex CLI e permissoes do ambiente.", + ), + ) from exc + + stdout = completed.stdout.strip() + stderr = completed.stderr.strip() + if completed.returncode != 0: + raise HTTPException( + status_code=502, + detail=_api_error( + "Codex CLI falhou", + stderr or stdout or f"Processo retornou codigo {completed.returncode}.", + "Confira login do Codex, rede, profile/modelo e tente novamente.", + raw="\n".join(part for part in (stdout, stderr) if part), + ), + ) + + return { + "analysis": stdout, + "stderr": stderr, + "codex_bin": codex_bin, + } + + +def main() -> None: + parser = argparse.ArgumentParser(description="Local OCI Logging/Search viewer for TIA logs using OCI CLI.") + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", default=8765, type=int) + args = parser.parse_args() + + import uvicorn + + uvicorn.run("app.tools.oci_logs_viewer:app", host=args.host, port=args.port, reload=False) + + +if __name__ == "__main__": + main() diff --git a/src/app/tools/sofya_transcribe.py b/src/app/tools/sofya_transcribe.py new file mode 100644 index 0000000..22108e9 --- /dev/null +++ b/src/app/tools/sofya_transcribe.py @@ -0,0 +1,1617 @@ +from __future__ import annotations + +import argparse +import asyncio +import copy +import json +import os +import time +import wave +from dataclasses import asdict +from pathlib import Path +from typing import Any + +import httpx +from dotenv import load_dotenv + +from app.common.call_config import resolve_vad_logging_overrides, resolve_vad_overrides +from app.providers.stt_config import override_config +from app.providers.stt_internal_livekit import ( + InternalHTTPSTT, + InternalSTTConfig, + extract_text_from_transcript, +) +from app.tools.local_stresstest.audio import ( + AudioSample, + audio_metrics, + pcm_to_wav_bytes, + vad_proxy_metrics, + write_wav, +) + + +# Ajuste estes valores direto no arquivo se quiser rodar sem flags. +DEFAULT_ENV_FILE = ".env.kube.fqa" +DEFAULT_AUDIO_PATH = "" +DEFAULT_OUTPUT_JSON = "" +DEFAULT_SENT_WAV_PATH = ".run/sofya-transcribe/sent_to_sofya.wav" +DEFAULT_SEGMENTS_DIR = ".run/sofya-transcribe/livekit_vad_segments" +DEFAULT_TRANSCRIPTS_OUTPUT = ".run/sofya-transcribe/silero_segments_transcripts.md" +DEFAULT_INPUT_CHANNEL = "cliente" +DEFAULT_CLIENT_CHANNEL = 1 +DEFAULT_AGENT_CHANNEL = 2 +DEFAULT_USE_LIVEKIT_VAD = True +DEFAULT_VAD_ONLY = False +DEFAULT_VAD_FRAME_MS = 20 +DEFAULT_VAD_TAIL_SILENCE_MS = 1600 + +DEFAULT_CALL_CONFIG = { + "callConfig": { + "vad": { + "minSpeechDuration": 0.15, + "activationThreshold": 0.30, + "deactivationThreshold": 0.15, + "minSilenceDuration": 1.0, + "prefixPaddingDuration": 1.0, + }, + "vadLogging": { + "logDecisions": True, + "logActivity": True, + "activityMinProbability": 0.03, + }, + "ws": { + "outputGain": 2.0, + }, + } +} + +# VAD do Sofya enviado dentro de override_config.processor.config.extra_params. +# Use None para herdar do env STT_VAD_* ou do override_config padrao. +DEFAULT_SOFYA_VAD_FILTER: bool | None = None +DEFAULT_SOFYA_VAD_THRESHOLD: float | None = None +DEFAULT_SOFYA_VAD_MIN_SILENCE_MS: int | None = None +DEFAULT_SOFYA_VAD_SPEECH_PAD_MS: int | None = None + +# Diagnostico local por RMS. Nao substitui o VAD do Sofya, so ajuda a enxergar +# onde o inicio de fala pode estar sendo comido. +DEFAULT_VAD_PROXY_THRESHOLD_DBFS = -45.0 +DEFAULT_VAD_PROXY_PREFIX_PADDING_MS: int | None = None +DEFAULT_VAD_PROXY_MIN_SPEECH_MS = 100 + +# Padding extra antes do WAV enviado ao Sofya, para simular STT_INPUT_PREFIX_PADDING_MS. +DEFAULT_STT_INPUT_PREFIX_PADDING_MS: int | None = None +DEFAULT_STT_INPUT_SUFFIX_PADDING_MS: int | None = None +DEFAULT_STT_PLAYBACK_SPEED = 0.85 +DEFAULT_STT_VOLUME_GAIN_DB = 6.0 +AUDIO_EXTENSIONS = {".wav", ".mp3", ".flac", ".ogg", ".m4a"} + + +def _env_str(name: str, default: str = "") -> str: + value = os.getenv(name) + return str(default if value is None else value).strip() + + +def _env_float(name: str, default: float | None = None) -> float | None: + value = _env_str(name) + if not value: + return default + try: + return float(value) + except ValueError: + return default + + +def _env_int(name: str, default: int | None = None) -> int | None: + value = _env_str(name) + if not value: + return default + try: + return int(value) + except ValueError: + return default + + +def _env_bool(name: str, default: bool | None = None) -> bool | None: + value = _env_str(name) + if not value: + return default + return value.lower() in {"1", "true", "yes", "y", "on", "sim"} + + +def _coalesce(*values: Any) -> Any: + for value in values: + if value is not None and value != "": + return value + return None + + +def _json_dumps(payload: Any) -> str: + return json.dumps(payload, ensure_ascii=False, indent=2) + + +def _float_override( + value: Any, + default: float, + *, + min_value: float | None = None, + max_value: float | None = None, +) -> float: + if value in (None, ""): + return default + try: + resolved = float(value) + except (TypeError, ValueError): + return default + if min_value is not None and resolved < min_value: + return default + if max_value is not None and resolved > max_value: + return default + return resolved + + +def _bool_override(value: Any, default: bool) -> bool: + if value in (None, ""): + return default + raw = str(value).strip().lower() + if raw in {"1", "true", "yes", "on", "sim"}: + return True + if raw in {"0", "false", "no", "off", "nao", "não"}: + return False + return default + + +def _default_livekit_vad_config() -> dict[str, float]: + activation_threshold = float(_env_float("VAD_ACTIVATION_THRESHOLD", 0.35) or 0.35) + prefix_padding_duration = float(_env_float("VAD_PREFIX_PADDING_DURATION", 1.5) or 1.5) + prefix_padding_min_duration = float(_env_float("VAD_PREFIX_PADDING_MIN_DURATION", 1.5) or 1.5) + return { + "min_speech_duration": float(_env_float("VAD_MIN_SPEECH_DURATION", 0.15) or 0.15), + "min_silence_duration": float( + _coalesce( + _env_float("VAD_MIN_SILENCE_DURATION"), + _env_float("LIVEKIT_VAD_MIN_SILENCE_DURATION_S"), + 0.35, + ) + ), + "activation_threshold": activation_threshold, + "deactivation_threshold": float( + _env_float( + "VAD_DEACTIVATION_THRESHOLD", + max(activation_threshold - 0.15, 0.01), + ) + or max(activation_threshold - 0.15, 0.01) + ), + "prefix_padding_duration": max(prefix_padding_duration, prefix_padding_min_duration), + } + + +def _default_livekit_vad_logging_config() -> dict[str, Any]: + return { + "log_decisions": _env_bool("FLOW_LOG_VAD_DECISIONS", False), + "log_activity": _env_bool("FLOW_LOG_VAD_ACTIVITY", False), + "activity_min_probability": float(_env_float("FLOW_LOG_VAD_ACTIVITY_MIN_PROB", 0.03) or 0.03), + } + + +def _call_config_payload(args: argparse.Namespace) -> dict[str, Any]: + if args.call_config_json: + payload = json.loads(args.call_config_json) + elif args.call_config_file: + payload = json.loads(Path(args.call_config_file).expanduser().read_text(encoding="utf-8")) + else: + payload = copy.deepcopy(DEFAULT_CALL_CONFIG) + if not isinstance(payload, dict): + raise ValueError("callConfig precisa ser um objeto JSON") + call_config = payload.get("callConfig", payload) + if not isinstance(call_config, dict): + raise ValueError("callConfig precisa ser um objeto JSON") + call_config = copy.deepcopy(call_config) + vad = call_config.setdefault("vad", {}) + if args.lk_min_speech_duration is not None: + vad["minSpeechDuration"] = args.lk_min_speech_duration + if args.lk_activation_threshold is not None: + vad["activationThreshold"] = args.lk_activation_threshold + if args.lk_deactivation_threshold is not None: + vad["deactivationThreshold"] = args.lk_deactivation_threshold + if args.lk_min_silence_duration is not None: + vad["minSilenceDuration"] = args.lk_min_silence_duration + if args.lk_prefix_padding_duration is not None: + vad["prefixPaddingDuration"] = args.lk_prefix_padding_duration + + vad_logging = call_config.setdefault("vadLogging", {}) + if args.lk_log_activity is not None: + vad_logging["logActivity"] = bool(args.lk_log_activity) + if args.lk_activity_min_probability is not None: + vad_logging["activityMinProbability"] = args.lk_activity_min_probability + return call_config + + +def _resolve_livekit_vad_config(args: argparse.Namespace) -> tuple[dict[str, float], dict[str, Any], dict[str, Any]]: + call_config = _call_config_payload(args) + defaults = _default_livekit_vad_config() + overrides = resolve_vad_overrides(call_config) + prefix_padding_min_duration = float(_env_float("VAD_PREFIX_PADDING_MIN_DURATION", 0.0) or 0.0) + vad_config = { + "min_speech_duration": _float_override( + overrides.get("min_speech_duration"), + defaults["min_speech_duration"], + min_value=0.0, + ), + "min_silence_duration": _float_override( + overrides.get("min_silence_duration"), + defaults["min_silence_duration"], + min_value=0.0, + ), + "activation_threshold": _float_override( + overrides.get("activation_threshold"), + defaults["activation_threshold"], + min_value=0.0, + max_value=1.0, + ), + "deactivation_threshold": _float_override( + overrides.get("deactivation_threshold"), + defaults["deactivation_threshold"], + min_value=0.0, + max_value=1.0, + ), + "prefix_padding_duration": _float_override( + overrides.get("prefix_padding_duration"), + defaults["prefix_padding_duration"], + min_value=prefix_padding_min_duration, + ), + } + + logging_defaults = _default_livekit_vad_logging_config() + logging_overrides = resolve_vad_logging_overrides(call_config) + logging_config = { + "log_decisions": _bool_override( + logging_overrides.get("log_decisions"), + bool(logging_defaults["log_decisions"]), + ), + "log_activity": _bool_override( + logging_overrides.get("log_activity"), + bool(logging_defaults["log_activity"]), + ), + "activity_min_probability": _float_override( + logging_overrides.get("activity_min_probability"), + float(logging_defaults["activity_min_probability"]), + min_value=0.0, + max_value=1.0, + ), + } + return vad_config, logging_config, call_config + + +def _build_arg_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description="Transcreve um arquivo de audio usando o STT Sofya HTTP do projeto.", + ) + parser.add_argument( + "--env-file", + default=DEFAULT_ENV_FILE, + help=f"Arquivo .env para carregar antes de chamar o Sofya. Default: {DEFAULT_ENV_FILE}", + ) + parser.add_argument( + "--override-env", + action="store_true", + help="Faz o .env sobrescrever variaveis ja existentes no processo.", + ) + parser.add_argument( + "--audio", + default=DEFAULT_AUDIO_PATH, + help="Caminho do arquivo ou diretorio de audios. Tambem pode usar SOFYA_AUDIO_FILE.", + ) + parser.add_argument( + "--audio-glob", + default="*.wav", + help="Padrao usado quando --audio aponta para diretorio. Ex.: '*.wav' ou '*.*'.", + ) + parser.add_argument( + "--input-channel", + default=DEFAULT_INPUT_CHANNEL, + help=( + "Canal do WAV usado antes do VAD/STT: mix, cliente/client/left/1 ou " + "agente/agent/right/2. Default: cliente." + ), + ) + mode_group = parser.add_mutually_exclusive_group() + mode_group.add_argument( + "--livekit-vad", + dest="use_livekit_vad", + action="store_true", + default=DEFAULT_USE_LIVEKIT_VAD, + help="Roda o audio pelo Silero VAD do LiveKit antes de chamar o Sofya. Default.", + ) + mode_group.add_argument( + "--direct-sofya", + dest="use_livekit_vad", + action="store_false", + help="Chama o Sofya direto com o audio inteiro, sem Silero VAD do LiveKit.", + ) + parser.add_argument( + "--vad-only", + action="store_true", + default=DEFAULT_VAD_ONLY, + help="Roda apenas o Silero VAD e salva os segmentos, sem chamar o Sofya.", + ) + parser.add_argument( + "--call-config-json", + default="", + help="JSON com callConfig para resolver os parametros de VAD do LiveKit.", + ) + parser.add_argument( + "--call-config-file", + default="", + help="Arquivo JSON com callConfig para resolver os parametros de VAD do LiveKit.", + ) + parser.add_argument("--url", default="", help="Override de STT_URL.") + parser.add_argument("--api-key", default="", help="Override de STT_KEY/STT_API_KEY.") + parser.add_argument("--language", default="", help="Override de STT_LANG.") + parser.add_argument("--timeout-s", type=float, default=None, help="Timeout HTTP do Sofya.") + parser.add_argument( + "--output-mode", + choices=["threshold_text", "api_text", "raw_json"], + default="", + help="Modo de saida do provider. Default vem de STT_OUTPUT_MODE ou threshold_text.", + ) + parser.add_argument( + "--min-prob-single-word", + type=float, + default=None, + help="Filtro local para palavra unica no modo threshold_text.", + ) + parser.add_argument( + "--initial-prompt", + default="", + help="Initial prompt enviado no override_config do Sofya.", + ) + + livekit_vad_group = parser.add_argument_group("Silero VAD do LiveKit") + livekit_vad_group.add_argument( + "--lk-min-speech-duration", + type=float, + default=None, + help="callConfig.vad.minSpeechDuration, em segundos.", + ) + livekit_vad_group.add_argument( + "--lk-activation-threshold", + type=float, + default=None, + help="callConfig.vad.activationThreshold.", + ) + livekit_vad_group.add_argument( + "--lk-deactivation-threshold", + type=float, + default=None, + help="callConfig.vad.deactivationThreshold.", + ) + livekit_vad_group.add_argument( + "--lk-min-silence-duration", + type=float, + default=None, + help="callConfig.vad.minSilenceDuration, em segundos.", + ) + livekit_vad_group.add_argument( + "--lk-prefix-padding-duration", + type=float, + default=None, + help="callConfig.vad.prefixPaddingDuration, em segundos.", + ) + livekit_vad_logging = livekit_vad_group.add_mutually_exclusive_group() + livekit_vad_logging.add_argument( + "--lk-log-activity", + dest="lk_log_activity", + action="store_true", + default=None, + help="Habilita logs/metricas de atividade do VAD no resumo.", + ) + livekit_vad_logging.add_argument( + "--no-lk-log-activity", + dest="lk_log_activity", + action="store_false", + help="Desabilita logs/metricas de atividade do VAD no resumo.", + ) + livekit_vad_group.add_argument( + "--lk-activity-min-probability", + type=float, + default=None, + help="callConfig.vadLogging.activityMinProbability.", + ) + livekit_vad_group.add_argument( + "--vad-frame-ms", + type=int, + default=DEFAULT_VAD_FRAME_MS, + help="Tamanho dos frames empurrados para o LiveKit VAD.", + ) + livekit_vad_group.add_argument( + "--vad-tail-silence-ms", + type=int, + default=DEFAULT_VAD_TAIL_SILENCE_MS, + help="Silencio ao final para forcar END_OF_SPEECH em audio sem cauda.", + ) + livekit_vad_group.add_argument( + "--segments-dir", + default=DEFAULT_SEGMENTS_DIR, + help="Diretorio onde os trechos cortados pelo LiveKit VAD serao salvos.", + ) + livekit_vad_group.add_argument( + "--transcripts-output", + default=DEFAULT_TRANSCRIPTS_OUTPUT, + help="Markdown com segmento Silero, WAV enviado ao Sofya e transcricao por segmento.", + ) + + vad_group = parser.add_argument_group("VAD interno do Sofya, nao e o Silero do LiveKit") + vad_filter = vad_group.add_mutually_exclusive_group() + vad_filter.add_argument( + "--vad-filter", + dest="vad_filter", + action="store_true", + default=DEFAULT_SOFYA_VAD_FILTER, + help="Envia vad_filter=true no extra_params do Sofya.", + ) + vad_filter.add_argument( + "--no-vad-filter", + dest="vad_filter", + action="store_false", + help="Envia vad_filter=false no extra_params do Sofya.", + ) + vad_group.add_argument( + "--vad-threshold", + type=float, + default=DEFAULT_SOFYA_VAD_THRESHOLD, + help="threshold em vad_parameters do Sofya. Ex.: 0.72", + ) + vad_group.add_argument( + "--vad-min-silence-ms", + type=int, + default=DEFAULT_SOFYA_VAD_MIN_SILENCE_MS, + help="min_silence_duration_ms em vad_parameters do Sofya.", + ) + vad_group.add_argument( + "--vad-speech-pad-ms", + type=int, + default=DEFAULT_SOFYA_VAD_SPEECH_PAD_MS, + help="speech_pad_ms em vad_parameters do Sofya.", + ) + vad_group.add_argument( + "--stt-input-prefix-padding-ms", + type=int, + default=DEFAULT_STT_INPUT_PREFIX_PADDING_MS, + help="Silencio adicionado antes do WAV enviado ao Sofya.", + ) + vad_group.add_argument( + "--stt-input-suffix-padding-ms", + type=int, + default=DEFAULT_STT_INPUT_SUFFIX_PADDING_MS, + help="Silencio adicionado no fim do WAV enviado ao Sofya.", + ) + vad_group.add_argument( + "--stt-playback-speed", + type=float, + default=None, + help=( + "Velocidade do audio enviado ao Sofya. 1.0 mantem original; " + "0.85 envia mais lento; 1.15 envia mais rapido." + ), + ) + vad_group.add_argument( + "--stt-volume-gain-db", + type=float, + default=None, + help=( + "Ganho em dB aplicado no audio enviado ao Sofya. " + "0 mantem original; 6 dobra aproximadamente a amplitude." + ), + ) + + proxy_group = parser.add_argument_group("Diagnostico VAD proxy local") + proxy_group.add_argument( + "--vad-proxy-threshold-dbfs", + type=float, + default=DEFAULT_VAD_PROXY_THRESHOLD_DBFS, + help="Limiar RMS local para diagnostico. Default: -45", + ) + proxy_group.add_argument( + "--vad-proxy-prefix-padding-ms", + type=int, + default=DEFAULT_VAD_PROXY_PREFIX_PADDING_MS, + help="Padding local do diagnostico. Default vem de VAD_PREFIX_PADDING_DURATION.", + ) + proxy_group.add_argument( + "--vad-proxy-min-speech-ms", + type=int, + default=DEFAULT_VAD_PROXY_MIN_SPEECH_MS, + help="Fala minima usada no diagnostico local.", + ) + + parser.add_argument( + "--save-sent-wav", + default=DEFAULT_SENT_WAV_PATH, + help="Salva o WAV normalizado que foi enviado ao Sofya. Use vazio para desativar.", + ) + parser.add_argument( + "--json-output", + default=DEFAULT_OUTPUT_JSON, + help="Opcional: salva um resumo JSON da execucao.", + ) + parser.add_argument( + "--print-override-config", + action="store_true", + help="Imprime o override_config final enviado ao Sofya.", + ) + return parser + + +def _audio_path_from_args(args: argparse.Namespace) -> Path: + raw_path = _coalesce(args.audio, _env_str("SOFYA_AUDIO_FILE"), _env_str("STRESS_AUDIO")) + if not raw_path: + raise RuntimeError( + "Informe o audio com --audio caminho.wav, SOFYA_AUDIO_FILE ou editando DEFAULT_AUDIO_PATH." + ) + return Path(str(raw_path)).expanduser() + + +def _resolve_audio_paths(args: argparse.Namespace) -> list[Path]: + root = _audio_path_from_args(args) + if root.is_dir(): + paths = [ + path + for path in sorted(root.glob(str(args.audio_glob or "*.wav"))) + if path.is_file() and path.suffix.lower() in AUDIO_EXTENSIONS + ] + if not paths: + raise RuntimeError(f"Nenhum audio encontrado em {root} com padrao {args.audio_glob!r}.") + return paths + return [root] + + +def _safe_stem(path: Path) -> str: + return "".join(ch if ch.isalnum() or ch in {"-", "_"} else "_" for ch in path.stem).strip("_") or "audio" + + +def _channel_spec_from_args(args: argparse.Namespace) -> str: + return str(_coalesce(args.input_channel, _env_str("SOFYA_INPUT_CHANNEL"), DEFAULT_INPUT_CHANNEL)).strip() + + +def _resolve_channel_index(channel_spec: str, channels: int) -> int | None: + spec = (channel_spec or DEFAULT_INPUT_CHANNEL).strip().lower() + if spec in {"mix", "mixed", "mono", "ambos", "todos", "all"}: + return None + if spec in {"cliente", "client", "customer", "left", "l"}: + channel_number = DEFAULT_CLIENT_CHANNEL if spec in {"cliente", "client", "customer"} else 1 + elif spec in {"agente", "agent", "right", "r"}: + channel_number = DEFAULT_AGENT_CHANNEL if spec in {"agente", "agent"} else 2 + else: + try: + channel_number = int(spec) + except ValueError as exc: + raise ValueError( + f"Canal invalido: {channel_spec}. Use mix, cliente, agente, left, right, 1 ou 2." + ) from exc + + if channel_number < 1 or channel_number > channels: + raise ValueError(f"Canal {channel_number} indisponivel: o audio tem {channels} canal(is).") + return channel_number - 1 + + +def _extract_pcm_channel(pcm: bytes, *, channels: int, channel_index: int) -> bytes: + if channels == 1: + return pcm + sample_width = 2 + frame_width = sample_width * channels + offset = channel_index * sample_width + out = bytearray() + for pos in range(0, len(pcm) - frame_width + 1, frame_width): + out.extend(pcm[pos + offset : pos + offset + sample_width]) + return bytes(out) + + +def _mix_pcm_channels(pcm: bytes, *, channels: int) -> bytes: + import audioop + + if channels == 1: + return pcm + if channels == 2: + return audioop.tomono(pcm, 2, 0.5, 0.5) + + sample_width = 2 + frame_width = sample_width * channels + out = bytearray() + for pos in range(0, len(pcm) - frame_width + 1, frame_width): + total = 0 + for channel_index in range(channels): + start = pos + (channel_index * sample_width) + total += int.from_bytes(pcm[start : start + sample_width], byteorder="little", signed=True) + average = max(-32768, min(32767, round(total / channels))) + out.extend(int(average).to_bytes(sample_width, byteorder="little", signed=True)) + return bytes(out) + + +def _resample_mono16(pcm: bytes, *, sample_rate: int, target_sample_rate: int = 16000) -> bytes: + import audioop + + if sample_rate == target_sample_rate: + return pcm + converted, _ = audioop.ratecv(pcm, 2, 1, sample_rate, target_sample_rate, None) + return converted + + +def _load_audio(path: Path, *, input_channel: str) -> AudioSample: + if not path.exists(): + raise FileNotFoundError(f"Audio nao encontrado: {path}") + try: + return _load_wav_audio(path, input_channel=input_channel) + except Exception as wav_error: + try: + return _load_with_soundfile(path, input_channel=input_channel) + except Exception as soundfile_error: + raise ValueError( + f"Nao consegui ler {path}. Esperado WAV PCM16, ou formato aceito pelo soundfile. " + f"wave={wav_error}; soundfile={soundfile_error}" + ) from soundfile_error + + +def _load_wav_audio(path: Path, *, input_channel: str) -> AudioSample: + with wave.open(str(path), "rb") as handle: + channels = handle.getnchannels() + sample_width = handle.getsampwidth() + sample_rate = handle.getframerate() + pcm = handle.readframes(handle.getnframes()) + + if sample_width != 2: + raise ValueError(f"{path} precisa ser PCM16 para leitura via wave, sample_width={sample_width}") + + channel_index = _resolve_channel_index(input_channel, channels) + if channel_index is None: + mono_pcm = _mix_pcm_channels(pcm, channels=channels) + channel_name = "mix" + else: + mono_pcm = _extract_pcm_channel(pcm, channels=channels, channel_index=channel_index) + channel_name = f"ch{channel_index + 1}" + mono_pcm = _resample_mono16(mono_pcm, sample_rate=sample_rate) + return AudioSample(name=f"{path.stem}_{channel_name}", pcm=mono_pcm, sample_rate=16000, channels=1) + + +def _load_with_soundfile(path: Path, *, input_channel: str) -> AudioSample: + import audioop + + import soundfile as sf + + data, sample_rate = sf.read(str(path), dtype="int16", always_2d=True) + if data.size <= 0: + raise ValueError(f"Audio vazio: {path}") + channels = int(data.shape[1]) + channel_index = _resolve_channel_index(input_channel, channels) + if channel_index is None: + data = data.mean(axis=1).astype("int16").reshape(-1, 1) + channel_name = "mix" + else: + data = data[:, channel_index].astype("int16").reshape(-1, 1) + channel_name = f"ch{channel_index + 1}" + pcm = data.tobytes() + if int(sample_rate) != 16000: + pcm, _ = audioop.ratecv(pcm, 2, 1, int(sample_rate), 16000, None) + sample_rate = 16000 + return AudioSample(name=f"{path.stem}_{channel_name}", pcm=pcm, sample_rate=int(sample_rate), channels=1) + + +def build_sofya_override_config(args: argparse.Namespace) -> dict[str, Any]: + cfg = copy.deepcopy(override_config) + processor = cfg.setdefault("processor", {}).setdefault("config", {}) + extra_params = processor.setdefault("extra_params", {}) + vad_params = extra_params.setdefault("vad_parameters", {}) + + vad_filter = _coalesce(args.vad_filter, _env_bool("STT_VAD_FILTER")) + if vad_filter is not None: + extra_params["vad_filter"] = bool(vad_filter) + + vad_threshold = _coalesce(args.vad_threshold, _env_float("STT_VAD_THRESHOLD")) + if vad_threshold is not None: + vad_params["threshold"] = float(vad_threshold) + + vad_min_silence_ms = _coalesce( + args.vad_min_silence_ms, + _env_int("STT_VAD_MIN_SILENCE_MS"), + ) + if vad_min_silence_ms is not None: + vad_params["min_silence_duration_ms"] = max(0, int(vad_min_silence_ms)) + + vad_speech_pad_ms = _coalesce( + args.vad_speech_pad_ms, + _env_int("STT_VAD_SPEECH_PAD_MS"), + ) + if vad_speech_pad_ms is not None: + vad_params["speech_pad_ms"] = max(0, int(vad_speech_pad_ms)) + + prompt = _coalesce( + args.initial_prompt, + _env_str("STT_FORCE_INITIAL_PROMPT"), + _env_str("STT_INITIAL_PROMPT"), + ) + if prompt: + extra_params["initial_prompt"] = str(prompt) + + return cfg + + +def _stt_config(args: argparse.Namespace, config_override: dict[str, Any]) -> InternalSTTConfig: + url = _coalesce(args.url, _env_str("STT_URL")) + if not url: + raise RuntimeError("STT_URL nao configurado. Use --url ou carregue um --env-file com STT_URL.") + + api_key = _coalesce(args.api_key, _env_str("STT_KEY"), _env_str("STT_API_KEY"), "unknow") + timeout_s = _coalesce(args.timeout_s, _env_float("STT_TIMEOUT_S"), 30.0) + min_prob_single_word = _coalesce( + args.min_prob_single_word, + _env_float("STT_MIN_PROB_SINGLE_WORD"), + 0.10, + ) + output_mode = _coalesce(args.output_mode, _env_str("STT_OUTPUT_MODE"), "threshold_text") + return InternalSTTConfig( + url=str(url), + api_key=str(api_key), + language=str(_coalesce(args.language, _env_str("STT_LANG"), "portuguese")), + timeout_s=float(timeout_s), + min_prob_single_word=float(min_prob_single_word), + config_override=json.dumps(config_override, ensure_ascii=False), + output_mode=str(output_mode), + ) + + +def _apply_prefix_padding(sample: AudioSample, padding_ms: int) -> tuple[AudioSample, int]: + padding_ms = max(0, int(padding_ms)) + if padding_ms <= 0: + return sample, 0 + frames = round(sample.sample_rate * padding_ms / 1000) + silence = b"\x00" * frames * sample.channels * 2 + if not silence: + return sample, 0 + actual_ms = round((frames / sample.sample_rate) * 1000) + return ( + AudioSample( + name=f"{sample.name}_stt_padded", + pcm=silence + sample.pcm, + sample_rate=sample.sample_rate, + channels=sample.channels, + ), + actual_ms, + ) + + +def _apply_suffix_padding(sample: AudioSample, padding_ms: int) -> tuple[AudioSample, int]: + padding_ms = max(0, int(padding_ms)) + if padding_ms <= 0: + return sample, 0 + frames = round(sample.sample_rate * padding_ms / 1000) + silence = b"\x00" * frames * sample.channels * 2 + if not silence: + return sample, 0 + actual_ms = round((frames / sample.sample_rate) * 1000) + return ( + AudioSample( + name=f"{sample.name}_stt_tail_padded", + pcm=sample.pcm + silence, + sample_rate=sample.sample_rate, + channels=sample.channels, + ), + actual_ms, + ) + + +def _apply_playback_speed(sample: AudioSample, speed: float) -> AudioSample: + speed = float(speed) + if speed <= 0: + raise ValueError("--stt-playback-speed precisa ser maior que zero") + if abs(speed - 1.0) < 0.0001: + return sample + + import audioop + + # Resample para outro numero de amostras e mantem o sample_rate original. + # Ex.: speed=0.85 gera mais amostras e o WAV toca mais lento. + resample_rate = max(1, round(sample.sample_rate / speed)) + pcm, _ = audioop.ratecv( + sample.pcm, + 2, + sample.channels, + sample.sample_rate, + resample_rate, + None, + ) + speed_name = f"{speed:.3f}".rstrip("0").rstrip(".").replace(".", "p") + return AudioSample( + name=f"{sample.name}_speed_{speed_name}x", + pcm=pcm, + sample_rate=sample.sample_rate, + channels=sample.channels, + ) + + +def _apply_volume_gain_db(sample: AudioSample, gain_db: float) -> AudioSample: + gain_db = float(gain_db) + if abs(gain_db) < 0.0001: + return sample + + import audioop + + factor = 10 ** (gain_db / 20.0) + pcm = audioop.mul(sample.pcm, 2, factor) + gain_name = ( + f"{gain_db:+.1f}db" + .replace("+", "plus") + .replace("-", "minus") + .replace(".", "p") + ) + return AudioSample( + name=f"{sample.name}_gain_{gain_name}", + pcm=pcm, + sample_rate=sample.sample_rate, + channels=sample.channels, + ) + + +def _pcm_duration_ms(sample: AudioSample) -> int: + if not sample.pcm: + return 0 + frames = len(sample.pcm) // (2 * max(1, sample.channels)) + return round((frames / max(1, sample.sample_rate)) * 1000) + + +def _tail_silence(sample: AudioSample, duration_ms: int) -> bytes: + frame_count = round(sample.sample_rate * max(0, duration_ms) / 1000) + return b"\x00" * frame_count * max(1, sample.channels) * 2 + + +def _iter_rtc_frames(sample: AudioSample, *, frame_ms: int, tail_silence_ms: int): + from livekit import rtc + + frame_ms = max(1, int(frame_ms)) + samples_per_frame = max(1, round(sample.sample_rate * frame_ms / 1000)) + bytes_per_frame = samples_per_frame * max(1, sample.channels) * 2 + data = sample.pcm + _tail_silence(sample, tail_silence_ms) + for offset in range(0, len(data), bytes_per_frame): + chunk = data[offset : offset + bytes_per_frame] + if not chunk: + continue + if len(chunk) < bytes_per_frame: + chunk += b"\x00" * (bytes_per_frame - len(chunk)) + yield rtc.AudioFrame( + data=chunk, + sample_rate=sample.sample_rate, + num_channels=sample.channels, + samples_per_channel=len(chunk) // (2 * max(1, sample.channels)), + ) + + +async def _run_livekit_vad( + sample: AudioSample, + *, + vad_config: dict[str, float], + logging_config: dict[str, Any], + frame_ms: int, + tail_silence_ms: int, + segments_dir: Path, +) -> tuple[list[dict[str, Any]], dict[str, Any]]: + from livekit.agents import vad as agents_vad + from livekit.plugins import silero + + segments_dir.mkdir(parents=True, exist_ok=True) + vad = silero.VAD.load( + min_speech_duration=vad_config["min_speech_duration"], + min_silence_duration=vad_config["min_silence_duration"], + activation_threshold=vad_config["activation_threshold"], + deactivation_threshold=vad_config["deactivation_threshold"], + prefix_padding_duration=vad_config["prefix_padding_duration"], + sample_rate=16000, + ) + stream = vad.stream() + segments: list[dict[str, Any]] = [] + diagnostics: dict[str, Any] = { + "inference_count": 0, + "max_probability": 0.0, + "activity_events": [], + "decision_events": [], + } + + async def _consume_events() -> None: + async for ev in stream: + if ev.type == agents_vad.VADEventType.INFERENCE_DONE: + probability = float(getattr(ev, "probability", 0.0) or 0.0) + diagnostics["inference_count"] += 1 + diagnostics["max_probability"] = max(float(diagnostics["max_probability"]), probability) + if ( + logging_config["log_activity"] + and probability >= logging_config["activity_min_probability"] + ): + diagnostics["activity_events"].append( + { + "timestamp_ms": round(float(ev.timestamp or 0.0) * 1000), + "probability": round(probability, 4), + "speaking": bool(getattr(ev, "speaking", False)), + "raw_speech_ms": round(float(getattr(ev, "raw_accumulated_speech", 0.0) or 0.0) * 1000), + "raw_silence_ms": round(float(getattr(ev, "raw_accumulated_silence", 0.0) or 0.0) * 1000), + } + ) + continue + + if ev.type == agents_vad.VADEventType.START_OF_SPEECH: + if logging_config["log_decisions"]: + diagnostics["decision_events"].append( + { + "type": "start_of_speech", + "timestamp_ms": round(float(ev.timestamp or 0.0) * 1000), + "speech_duration_ms": round(float(ev.speech_duration or 0.0) * 1000), + } + ) + continue + + if ev.type != agents_vad.VADEventType.END_OF_SPEECH: + continue + + pcm = b"".join(bytes(frame.data) for frame in ev.frames) + segment = AudioSample( + name=f"{sample.name}_segment_{len(segments) + 1:02d}", + pcm=pcm, + sample_rate=sample.sample_rate, + channels=sample.channels, + ) + duration_ms = _pcm_duration_ms(segment) + end_ms = round(float(ev.timestamp or 0.0) * 1000) + start_ms = max(0, end_ms - duration_ms) + segment_path = write_wav( + segments_dir / f"{len(segments) + 1:02d}_{segment.name}.wav", + segment, + ) + metrics = audio_metrics(segment.pcm, sample_rate=segment.sample_rate, channels=segment.channels) + row = { + "index": len(segments) + 1, + "path": str(segment_path), + "start_ms": start_ms, + "end_ms": end_ms, + "duration_ms": duration_ms, + "speech_duration_ms": round(float(ev.speech_duration or 0.0) * 1000), + "silence_duration_ms": round(float(ev.silence_duration or 0.0) * 1000), + "rms_dbfs": metrics.rms_dbfs, + "peak_dbfs": metrics.peak_dbfs, + "pcm": segment.pcm, + } + segments.append(row) + if logging_config["log_decisions"]: + diagnostics["decision_events"].append( + { + "type": "end_of_speech", + "timestamp_ms": end_ms, + "speech_duration_ms": row["speech_duration_ms"], + "silence_duration_ms": row["silence_duration_ms"], + "segment_duration_ms": duration_ms, + } + ) + + consumer = asyncio.create_task(_consume_events()) + try: + for frame in _iter_rtc_frames( + sample, + frame_ms=frame_ms, + tail_silence_ms=tail_silence_ms, + ): + stream.push_frame(frame) + await asyncio.sleep(0) + stream.end_input() + await consumer + finally: + if not consumer.done(): + consumer.cancel() + await asyncio.gather(consumer, return_exceptions=True) + try: + await stream.aclose() + except Exception: + pass + + return segments, diagnostics + + +def _resolve_vad_proxy_prefix_padding_ms(args: argparse.Namespace) -> int: + explicit = _coalesce( + args.vad_proxy_prefix_padding_ms, + _env_int("STRESS_VAD_PREFIX_PADDING_MS"), + ) + if explicit is not None: + return max(0, int(explicit)) + return max(0, round(float(_env_float("VAD_PREFIX_PADDING_DURATION", 1.0) or 1.0) * 1000)) + + +def _resolve_stt_input_padding_ms(args: argparse.Namespace) -> int: + value = _coalesce( + args.stt_input_prefix_padding_ms, + _env_int("STT_INPUT_PREFIX_PADDING_MS"), + 250, + ) + return max(0, int(value)) + + +def _resolve_stt_input_suffix_padding_ms(args: argparse.Namespace) -> int: + value = _coalesce( + args.stt_input_suffix_padding_ms, + _env_int("STT_INPUT_SUFFIX_PADDING_MS"), + 0, + ) + return max(0, int(value)) + + +def _resolve_stt_playback_speed(args: argparse.Namespace) -> float: + value = _coalesce( + args.stt_playback_speed, + _env_float("STT_PLAYBACK_SPEED"), + DEFAULT_STT_PLAYBACK_SPEED, + ) + speed = float(value) + if speed <= 0: + raise ValueError("--stt-playback-speed precisa ser maior que zero") + return speed + + +def _resolve_stt_volume_gain_db(args: argparse.Namespace) -> float: + value = _coalesce( + args.stt_volume_gain_db, + _env_float("STT_VOLUME_GAIN_DB"), + DEFAULT_STT_VOLUME_GAIN_DB, + ) + return float(value) + + +def _result_payload( + *, + audio_path: Path, + input_channel: str, + sent_wav_path: Path | None, + cfg: InternalSTTConfig, + override: dict[str, Any], + source_sample: AudioSample, + sent_sample: AudioSample, + applied_padding_ms: int, + applied_suffix_padding_ms: int, + stt_playback_speed: float, + stt_volume_gain_db: float, + vad_proxy: Any, + transcript_raw: str, + latency_ms: int, +) -> dict[str, Any]: + return { + "audio_path": str(audio_path), + "input_channel": input_channel, + "sent_wav_path": str(sent_wav_path) if sent_wav_path is not None else "", + "stt_url": cfg.url, + "language": cfg.language, + "output_mode": cfg.output_mode, + "latency_ms": latency_ms, + "source_audio": asdict(audio_metrics(source_sample.pcm, sample_rate=source_sample.sample_rate)), + "sent_audio": asdict(audio_metrics(sent_sample.pcm, sample_rate=sent_sample.sample_rate)), + "applied_stt_input_prefix_padding_ms": applied_padding_ms, + "applied_stt_input_suffix_padding_ms": applied_suffix_padding_ms, + "stt_playback_speed": stt_playback_speed, + "stt_volume_gain_db": stt_volume_gain_db, + "vad_proxy": asdict(vad_proxy), + "sofya_vad": override["processor"]["config"]["extra_params"].get("vad_parameters", {}), + "sofya_vad_filter": override["processor"]["config"]["extra_params"].get("vad_filter"), + "transcript": extract_text_from_transcript(transcript_raw), + "transcript_raw": transcript_raw, + } + + +async def _post_sample_to_sofya( + stt: InternalHTTPSTT, + *, + sample: AudioSample, + req_id: str, + language: str, +) -> dict[str, Any]: + metrics = audio_metrics(sample.pcm, sample_rate=sample.sample_rate, channels=sample.channels) + started = time.perf_counter() + event = await stt._post_internal_http( + sample.pcm, + pcm_to_wav_bytes(sample), + req_id=req_id, + started_ns=time.time_ns(), + language=language, + audio_duration_ms=metrics.duration_ms, + level_dbfs=metrics.rms_dbfs, + ) + latency_ms = round((time.perf_counter() - started) * 1000) + transcript_raw = event.alternatives[0].text if event.alternatives else "" + return { + "latency_ms": latency_ms, + "transcript_raw": transcript_raw, + "transcript": extract_text_from_transcript(transcript_raw), + } + + +def _print_livekit_vad_summary( + *, + vad_config: dict[str, float], + logging_config: dict[str, Any], + segments: list[dict[str, Any]], + diagnostics: dict[str, Any], +) -> None: + print("[sofya-transcribe] livekit_vad_config:") + print(_json_dumps(vad_config)) + print("[sofya-transcribe] livekit_vad_logging:") + print(_json_dumps(logging_config)) + print( + "[sofya-transcribe] livekit_vad " + f"segments={len(segments)} " + f"inference_count={diagnostics['inference_count']} " + f"max_probability={float(diagnostics['max_probability']):.4f}" + ) + for row in segments: + print( + "[sofya-transcribe] segment " + f"#{row['index']} start_ms={row['start_ms']} end_ms={row['end_ms']} " + f"duration_ms={row['duration_ms']} speech_ms={row['speech_duration_ms']} " + f"silence_ms={row['silence_duration_ms']} rms_dbfs={row['rms_dbfs']} " + f"path={row['path']}" + ) + + +def _format_ms(value: Any) -> str: + try: + total_ms = max(0, int(round(float(value)))) + except (TypeError, ValueError): + total_ms = 0 + minutes, rest = divmod(total_ms, 60_000) + seconds, ms = divmod(rest, 1000) + return f"{minutes:02d}:{seconds:02d}.{ms:03d}" + + +def _write_segment_transcripts( + path: Path, + *, + audio_path: Path, + input_channel: str, + vad_config: dict[str, float], + segments: list[dict[str, Any]], + transcript: str, +) -> Path: + path.parent.mkdir(parents=True, exist_ok=True) + lines = [ + "# Silero VAD -> Sofya STT", + "", + f"- Audio original: `{audio_path}`", + f"- Canal usado: `{input_channel}`", + "- Observacao: Silero nao transcreve texto; ele corta audio. Os textos abaixo sao retornos do Sofya para cada segmento cortado pelo Silero.", + "- VAD LiveKit:", + f" - min_speech_duration: `{vad_config['min_speech_duration']}`", + f" - activation_threshold: `{vad_config['activation_threshold']}`", + f" - deactivation_threshold: `{vad_config['deactivation_threshold']}`", + f" - min_silence_duration: `{vad_config['min_silence_duration']}`", + f" - prefix_padding_duration: `{vad_config['prefix_padding_duration']}`", + "", + "## Transcricao Final", + "", + transcript or "_Sem texto transcrito._", + "", + "## Segmentos", + "", + ] + for row in segments: + text = str(row.get("transcript") or "").strip() + lines.extend( + [ + f"### Segmento {int(row['index']):02d}", + "", + f"- Intervalo original: `{_format_ms(row.get('start_ms'))}` -> `{_format_ms(row.get('end_ms'))}`", + f"- Duracao enviada: `{row.get('duration_ms', '')} ms`", + f"- Fala estimada pelo VAD: `{row.get('speech_duration_ms', '')} ms`", + f"- Silencio de fechamento: `{row.get('silence_duration_ms', '')} ms`", + f"- RMS: `{row.get('rms_dbfs', '')} dBFS`", + f"- WAV cortado pelo Silero: `{row.get('path', '')}`", + f"- WAV enviado ao Sofya: `{row.get('sent_to_sofya_path', '')}`", + f"- Padding inicio enviado ao Sofya: `{row.get('applied_stt_input_prefix_padding_ms', '')} ms`", + f"- Padding fim enviado ao Sofya: `{row.get('applied_stt_input_suffix_padding_ms', '')} ms`", + f"- Velocidade enviada ao Sofya: `{row.get('stt_playback_speed', '')}`", + f"- Ganho enviado ao Sofya: `{row.get('stt_volume_gain_db', '')} dB`", + "", + "Transcricao:", + "", + text or "_Sem texto._", + "", + ] + ) + path.write_text("\n".join(lines).rstrip() + "\n", encoding="utf-8") + return path + + +def _write_batch_transcripts(path: Path, results: list[dict[str, Any]]) -> Path: + path.parent.mkdir(parents=True, exist_ok=True) + lines = [ + "# Batch Silero VAD -> Sofya STT", + "", + f"- Audios processados: `{len(results)}`", + "- Observacao: Silero corta audio; os textos sao retornos do Sofya para cada segmento enviado.", + "", + "## Resumo", + "", + "| audio | canal | segmentos | transcricao |", + "| --- | --- | ---: | --- |", + ] + for result in results: + transcript = str(result.get("transcript") or "").replace("\n", " ").strip() + if len(transcript) > 160: + transcript = transcript[:157].rstrip() + "..." + lines.append( + f"| `{result.get('audio_path', '')}` | `{result.get('input_channel', '')}` | " + f"{len(result.get('segments') or [])} | {transcript or '_Sem texto_'} |" + ) + + for audio_idx, result in enumerate(results, start=1): + lines.extend( + [ + "", + f"## Audio {audio_idx:02d}: `{result.get('audio_path', '')}`", + "", + f"- Canal usado: `{result.get('input_channel', '')}`", + f"- Segmentos: `{len(result.get('segments') or [])}`", + f"- Markdown individual: `{result.get('transcripts_output', '')}`", + "", + "Transcricao final:", + "", + str(result.get("transcript") or "").strip() or "_Sem texto transcrito._", + "", + "### Segmentos", + "", + ] + ) + for row in result.get("segments") or []: + text = str(row.get("transcript") or "").strip() + lines.extend( + [ + f"#### Segmento {int(row['index']):02d}", + "", + f"- Intervalo original: `{_format_ms(row.get('start_ms'))}` -> `{_format_ms(row.get('end_ms'))}`", + f"- Duracao enviada: `{row.get('duration_ms', '')} ms`", + f"- Fala estimada pelo VAD: `{row.get('speech_duration_ms', '')} ms`", + f"- Silencio de fechamento: `{row.get('silence_duration_ms', '')} ms`", + f"- RMS: `{row.get('rms_dbfs', '')} dBFS`", + f"- WAV cortado pelo Silero: `{row.get('path', '')}`", + f"- WAV enviado ao Sofya: `{row.get('sent_to_sofya_path', '')}`", + f"- Padding inicio enviado ao Sofya: `{row.get('applied_stt_input_prefix_padding_ms', '')} ms`", + f"- Padding fim enviado ao Sofya: `{row.get('applied_stt_input_suffix_padding_ms', '')} ms`", + f"- Velocidade enviada ao Sofya: `{row.get('stt_playback_speed', '')}`", + f"- Ganho enviado ao Sofya: `{row.get('stt_volume_gain_db', '')} dB`", + "", + "Transcricao:", + "", + text or "_Sem texto._", + "", + ] + ) + path.write_text("\n".join(lines).rstrip() + "\n", encoding="utf-8") + return path + + +async def _run_one_audio( + args: argparse.Namespace, + *, + audio_path: Path, + input_channel: str, + segments_dir: Path | None = None, + transcripts_output: Path | None = None, + json_output: Path | None = None, +) -> dict[str, Any]: + source_sample = _load_audio(audio_path, input_channel=input_channel) + source_metrics = audio_metrics( + source_sample.pcm, + sample_rate=source_sample.sample_rate, + channels=source_sample.channels, + ) + vad_proxy = vad_proxy_metrics( + source_sample, + threshold_dbfs=float(args.vad_proxy_threshold_dbfs), + prefix_padding_ms=_resolve_vad_proxy_prefix_padding_ms(args), + min_speech_ms=max(0, int(args.vad_proxy_min_speech_ms)), + ) + + print( + "[sofya-transcribe] audio=" + f"{audio_path} input_channel={input_channel} normalized_sample={source_sample.name} " + f"duration_ms={source_metrics.duration_ms} rms_dbfs={source_metrics.rms_dbfs}" + ) + print( + "[sofya-transcribe] vad_proxy " + f"first_voice_ms={vad_proxy.first_voice_ms} " + f"speech_ms={vad_proxy.speech_ms} " + f"low_start_risk={vad_proxy.low_start_risk}" + ) + stt_playback_speed = _resolve_stt_playback_speed(args) + stt_volume_gain_db = _resolve_stt_volume_gain_db(args) + print(f"[sofya-transcribe] stt_playback_speed={stt_playback_speed}") + print(f"[sofya-transcribe] stt_volume_gain_db={stt_volume_gain_db}") + + if args.use_livekit_vad: + vad_config, vad_logging_config, call_config = _resolve_livekit_vad_config(args) + segments_dir = segments_dir or Path(args.segments_dir).expanduser() + segments, diagnostics = await _run_livekit_vad( + source_sample, + vad_config=vad_config, + logging_config=vad_logging_config, + frame_ms=max(1, int(args.vad_frame_ms)), + tail_silence_ms=max(0, int(args.vad_tail_silence_ms)), + segments_dir=segments_dir, + ) + _print_livekit_vad_summary( + vad_config=vad_config, + logging_config=vad_logging_config, + segments=segments, + diagnostics=diagnostics, + ) + + if args.vad_only: + result = { + "mode": "livekit_vad_only", + "audio_path": str(audio_path), + "input_channel": input_channel, + "source_audio": asdict(source_metrics), + "vad_proxy": asdict(vad_proxy), + "call_config": call_config, + "livekit_vad": vad_config, + "livekit_vad_logging": vad_logging_config, + "livekit_vad_diagnostics": diagnostics, + "segments": [{key: value for key, value in row.items() if key != "pcm"} for row in segments], + } + if json_output is not None: + json_path = json_output + json_path.parent.mkdir(parents=True, exist_ok=True) + json_path.write_text(_json_dumps(result) + "\n", encoding="utf-8") + print(f"[sofya-transcribe] json_output={json_path}") + print("[sofya-transcribe] vad_only=true; Sofya nao foi chamado.") + return result + + config_override = build_sofya_override_config(args) + cfg = _stt_config(args, config_override) + os.environ.pop("STT_FORCE_OVERRIDE_CONFIG", None) + if args.print_override_config: + print("[sofya-transcribe] override_config enviado ao Sofya:") + print(_json_dumps(config_override)) + print( + "[sofya-transcribe] sofya_vad " + f"filter={config_override['processor']['config']['extra_params'].get('vad_filter')} " + f"params={config_override['processor']['config']['extra_params'].get('vad_parameters')}" + ) + + stt_input_padding_ms = _resolve_stt_input_padding_ms(args) + stt_input_suffix_padding_ms = _resolve_stt_input_suffix_padding_ms(args) + async with httpx.AsyncClient() as client: + stt = InternalHTTPSTT(cfg, client=client) + for row in segments: + segment_sample = AudioSample( + name=f"{source_sample.name}_segment_{row['index']:02d}", + pcm=row["pcm"], + sample_rate=source_sample.sample_rate, + channels=source_sample.channels, + ) + speed_adjusted_segment = _apply_playback_speed(segment_sample, stt_playback_speed) + gain_adjusted_segment = _apply_volume_gain_db( + speed_adjusted_segment, + stt_volume_gain_db, + ) + sent_segment, applied_padding_ms = _apply_prefix_padding( + gain_adjusted_segment, + stt_input_padding_ms, + ) + sent_segment, applied_suffix_padding_ms = _apply_suffix_padding( + sent_segment, + stt_input_suffix_padding_ms, + ) + sent_path = write_wav( + segments_dir / f"{row['index']:02d}_{sent_segment.name}_sent_to_sofya.wav", + sent_segment, + ) + stt_result = await _post_sample_to_sofya( + stt, + sample=sent_segment, + req_id=f"sofya-file-{audio_path.stem}-seg-{row['index']:02d}", + language=cfg.language, + ) + row.update( + { + "sent_to_sofya_path": str(sent_path), + "applied_stt_input_prefix_padding_ms": applied_padding_ms, + "applied_stt_input_suffix_padding_ms": applied_suffix_padding_ms, + "stt_playback_speed": stt_playback_speed, + "stt_volume_gain_db": stt_volume_gain_db, + "sofya_latency_ms": stt_result["latency_ms"], + "transcript": stt_result["transcript"], + "transcript_raw": stt_result["transcript_raw"], + } + ) + print( + "[sofya-transcribe] transcript " + f"segment=#{row['index']} latency_ms={row['sofya_latency_ms']} " + f"text={row['transcript']}" + ) + + transcript = " ".join(str(row.get("transcript") or "").strip() for row in segments).strip() + transcripts_output = transcripts_output if transcripts_output is not None else ( + Path(args.transcripts_output).expanduser() if args.transcripts_output else None + ) + if transcripts_output is not None: + transcripts_path = _write_segment_transcripts( + transcripts_output, + audio_path=audio_path, + input_channel=input_channel, + vad_config=vad_config, + segments=segments, + transcript=transcript, + ) + print(f"[sofya-transcribe] transcripts_output={transcripts_path}") + result = { + "mode": "livekit_vad_to_sofya", + "audio_path": str(audio_path), + "input_channel": input_channel, + "source_audio": asdict(source_metrics), + "vad_proxy": asdict(vad_proxy), + "call_config": call_config, + "livekit_vad": vad_config, + "livekit_vad_logging": vad_logging_config, + "livekit_vad_diagnostics": diagnostics, + "sofya_vad": config_override["processor"]["config"]["extra_params"].get("vad_parameters", {}), + "sofya_vad_filter": config_override["processor"]["config"]["extra_params"].get("vad_filter"), + "stt_input_prefix_padding_ms": stt_input_padding_ms, + "stt_input_suffix_padding_ms": stt_input_suffix_padding_ms, + "stt_playback_speed": stt_playback_speed, + "stt_volume_gain_db": stt_volume_gain_db, + "transcript": transcript, + "transcripts_output": str(transcripts_output) if transcripts_output is not None else "", + "segments": [{key: value for key, value in row.items() if key != "pcm"} for row in segments], + } + if json_output is not None: + json_path = json_output + json_path.parent.mkdir(parents=True, exist_ok=True) + json_path.write_text(_json_dumps(result) + "\n", encoding="utf-8") + print(f"[sofya-transcribe] json_output={json_path}") + print("[sofya-transcribe] transcript_final:") + print(transcript) + return result + + config_override = build_sofya_override_config(args) + cfg = _stt_config(args, config_override) + + # Garante que os parametros montados por este script ganhem de um + # STT_FORCE_OVERRIDE_CONFIG herdado do shell. + os.environ.pop("STT_FORCE_OVERRIDE_CONFIG", None) + + speed_adjusted_sample = _apply_playback_speed(source_sample, stt_playback_speed) + gain_adjusted_sample = _apply_volume_gain_db(speed_adjusted_sample, stt_volume_gain_db) + sent_sample, applied_padding_ms = _apply_prefix_padding( + gain_adjusted_sample, + _resolve_stt_input_padding_ms(args), + ) + sent_sample, applied_suffix_padding_ms = _apply_suffix_padding( + sent_sample, + _resolve_stt_input_suffix_padding_ms(args), + ) + sent_wav_path = Path(args.save_sent_wav).expanduser() if args.save_sent_wav else None + if sent_wav_path is not None: + write_wav(sent_wav_path, sent_sample) + + if args.print_override_config: + print("[sofya-transcribe] override_config enviado ao Sofya:") + print(_json_dumps(config_override)) + + print( + "[sofya-transcribe] sofya_vad " + f"filter={config_override['processor']['config']['extra_params'].get('vad_filter')} " + f"params={config_override['processor']['config']['extra_params'].get('vad_parameters')}" + ) + + async with httpx.AsyncClient() as client: + stt = InternalHTTPSTT(cfg, client=client) + stt_result = await _post_sample_to_sofya( + stt, + sample=sent_sample, + req_id=f"sofya-file-{audio_path.stem}", + language=cfg.language, + ) + latency_ms = stt_result["latency_ms"] + transcript_raw = stt_result["transcript_raw"] + transcript = extract_text_from_transcript(transcript_raw) + + result = _result_payload( + audio_path=audio_path, + input_channel=input_channel, + sent_wav_path=sent_wav_path, + cfg=cfg, + override=config_override, + source_sample=source_sample, + sent_sample=sent_sample, + applied_padding_ms=applied_padding_ms, + applied_suffix_padding_ms=applied_suffix_padding_ms, + stt_playback_speed=stt_playback_speed, + stt_volume_gain_db=stt_volume_gain_db, + vad_proxy=vad_proxy, + transcript_raw=transcript_raw, + latency_ms=latency_ms, + ) + if json_output is not None: + json_path = json_output + json_path.parent.mkdir(parents=True, exist_ok=True) + json_path.write_text(_json_dumps(result) + "\n", encoding="utf-8") + print(f"[sofya-transcribe] json_output={json_path}") + + print(f"[sofya-transcribe] latency_ms={latency_ms}") + print("[sofya-transcribe] transcript:") + print(transcript) + if cfg.output_mode == "raw_json": + print("[sofya-transcribe] transcript_raw:") + print(transcript_raw) + return result + + +async def run(args: argparse.Namespace) -> int: + env_file = Path(args.env_file).expanduser() if args.env_file else None + if env_file and env_file.exists(): + load_dotenv(env_file, override=bool(args.override_env)) + + input_channel = _channel_spec_from_args(args) + audio_paths = _resolve_audio_paths(args) + base_segments_dir = Path(args.segments_dir).expanduser() + base_transcripts_output = Path(args.transcripts_output).expanduser() if args.transcripts_output else None + base_json_output = Path(args.json_output).expanduser() if args.json_output else None + + if len(audio_paths) == 1: + await _run_one_audio( + args, + audio_path=audio_paths[0], + input_channel=input_channel, + segments_dir=base_segments_dir, + transcripts_output=base_transcripts_output, + json_output=base_json_output, + ) + return 0 + + print(f"[sofya-transcribe] batch audios={len(audio_paths)}") + results: list[dict[str, Any]] = [] + individual_dir = ( + base_transcripts_output.parent / f"{base_transcripts_output.stem}_files" + if base_transcripts_output is not None + else base_segments_dir / "reports" + ) + for idx, audio_path in enumerate(audio_paths, start=1): + safe_name = f"{idx:02d}_{_safe_stem(audio_path)}" + print(f"[sofya-transcribe] batch_item {idx}/{len(audio_paths)} audio={audio_path}") + result = await _run_one_audio( + args, + audio_path=audio_path, + input_channel=input_channel, + segments_dir=base_segments_dir / safe_name, + transcripts_output=individual_dir / f"{safe_name}_segments_transcripts.md", + json_output=None, + ) + results.append(result) + + if base_transcripts_output is not None: + batch_path = _write_batch_transcripts(base_transcripts_output, results) + print(f"[sofya-transcribe] batch_transcripts_output={batch_path}") + + if base_json_output is not None: + payload = { + "mode": "batch", + "input_channel": input_channel, + "audio_count": len(audio_paths), + "transcripts_output": str(base_transcripts_output) if base_transcripts_output is not None else "", + "results": results, + } + base_json_output.parent.mkdir(parents=True, exist_ok=True) + base_json_output.write_text(_json_dumps(payload) + "\n", encoding="utf-8") + print(f"[sofya-transcribe] batch_json_output={base_json_output}") + + return 0 + + +def main() -> None: + parser = _build_arg_parser() + args = parser.parse_args() + raise SystemExit(asyncio.run(run(args))) + + +if __name__ == "__main__": + main() diff --git a/src/app/tools/xai_tts_load.py b/src/app/tools/xai_tts_load.py new file mode 100644 index 0000000..8db450d --- /dev/null +++ b/src/app/tools/xai_tts_load.py @@ -0,0 +1,697 @@ +from __future__ import annotations + +import argparse +import asyncio +import base64 +import binascii +import csv +import json +import os +import time +import uuid +import wave +from collections import Counter +from dataclasses import asdict, dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from urllib.parse import urlencode + +import aiohttp +from dotenv import load_dotenv + + +DEFAULT_TEXT = ( + "Ola! Sou especialista em contas e estou a disposicao. " + "Como posso te ajudar com a sua fatura hoje?" +) +DEFAULT_OUTPUT_JSONL = Path(".run/xai-tts-load/results.jsonl") +DEFAULT_SUMMARY_CSV = Path(".run/xai-tts-load/summary.csv") +SAMPLE_RATE = 24_000 +NUM_CHANNELS = 1 +SAMPLE_WIDTH_BYTES = 2 + + +def _now_utc() -> str: + return datetime.now(timezone.utc).isoformat() + + +def _elapsed_ms(started_at: float) -> int: + return max(0, round((time.perf_counter() - started_at) * 1000)) + + +def _env_str(name: str, default: str = "") -> str: + return str(os.getenv(name, default) or default).strip() + + +def _bool_query(value: bool) -> str: + return "true" if value else "false" + + +def _percentile(values: list[int], percentile: float) -> int | None: + if not values: + return None + ordered = sorted(values) + index = round((len(ordered) - 1) * percentile) + return ordered[index] + + +def _write_wav(path: Path, pcm: bytes) -> Path: + path.parent.mkdir(parents=True, exist_ok=True) + with wave.open(str(path), "wb") as handle: + handle.setnchannels(NUM_CHANNELS) + handle.setsampwidth(SAMPLE_WIDTH_BYTES) + handle.setframerate(SAMPLE_RATE) + handle.writeframes(pcm) + return path + + +@dataclass(frozen=True, slots=True) +class Config: + env_file: Path + api_key: str + websocket_url: str + voice: str + language: str + text_normalization: bool + optimize_streaming_latency: int + iterations: int + concurrency: int + reuse_connection: bool + reset_on_failure: bool + prewarm_idle_s: float + idle_between_requests_s: float + request_timeout_s: float + connect_timeout_s: float + gap_timeout_s: float + output_jsonl: Path + summary_csv: Path | None + audio_dir: Path | None + texts: tuple[str, ...] + + +@dataclass(slots=True) +class XAIConnection: + connection_id: str + ws: aiohttp.ClientWebSocketResponse + opened_at: float + opened_connect_ms: int + turns_completed: int = 0 + last_outcome: str = "" + last_audio_bytes: int = 0 + last_audio_chunks: int = 0 + + @property + def closed(self) -> bool: + return self.ws.closed + + @property + def age_ms(self) -> int: + return _elapsed_ms(self.opened_at) + + +@dataclass(slots=True) +class Result: + run_id: str + request_id: str + worker_id: int + iteration: int + started_at: str + connection_id: str + connection_age_ms: int + connection_turn_index: int + reused_connection: bool + opened_for_request: bool + idle_before_turn_ms: int + previous_outcome: str + previous_audio_bytes: int + previous_audio_chunks: int + text_len: int + connect_ms: int + send_ms: int + first_audio_ms: int | None + request_to_first_audio_ms: int | None + total_to_first_audio_ms: int | None + elapsed_ms: int + audio_chunks: int + audio_bytes: int + empty_audio_deltas: int + max_gap_ms: int + gap_timeout_ms: int + outcome: str + trace_id: str + error: str + audio_path: str + + +def _build_url(config: Config) -> str: + params = { + "voice": config.voice, + "language": config.language, + "codec": "pcm", + "sample_rate": str(SAMPLE_RATE), + "optimize_streaming_latency": str(config.optimize_streaming_latency), + "text_normalization": _bool_query(config.text_normalization), + } + return f"{config.websocket_url}?{urlencode(params)}" + + +async def _open_connection( + session: aiohttp.ClientSession, + config: Config, + *, + worker_id: int, +) -> XAIConnection: + connection_id = f"direct-{worker_id}-{uuid.uuid4().hex[:8]}" + started_at = time.perf_counter() + ws = await asyncio.wait_for( + session.ws_connect( + _build_url(config), + headers={"Authorization": f"Bearer {config.api_key}"}, + heartbeat=20, + ), + timeout=config.connect_timeout_s, + ) + opened_at = time.perf_counter() + return XAIConnection( + connection_id=connection_id, + ws=ws, + opened_at=opened_at, + opened_connect_ms=max(0, round((opened_at - started_at) * 1000)), + ) + + +async def _close_connection(conn: XAIConnection | None, *, reason: str) -> None: + if conn is None or conn.closed: + return + await conn.ws.close(message=reason.encode("utf-8")) + + +async def _receive_xai_turn( + *, + config: Config, + conn: XAIConnection, + text: str, + run_id: str, + worker_id: int, + iteration: int, + job_started_at: float, + opened_for_request: bool, + idle_before_turn_ms: int, + output_audio_dir: Path | None, +) -> Result: + request_id = uuid.uuid4().hex[:12] + send_started_at = time.perf_counter() + await conn.ws.send_str(json.dumps({"type": "text.delta", "delta": text}, ensure_ascii=False)) + await conn.ws.send_str(json.dumps({"type": "text.done"})) + send_ms = _elapsed_ms(send_started_at) + + turn_started_at = time.perf_counter() + first_audio_ms: int | None = None + request_to_first_audio_ms: int | None = None + total_to_first_audio_ms: int | None = None + last_audio_at = 0.0 + audio_chunks = 0 + audio_bytes = 0 + empty_audio_deltas = 0 + max_gap_ms = 0 + trace_id = "" + error = "" + outcome = "unknown" + audio = bytearray() + previous_outcome = conn.last_outcome + previous_audio_bytes = conn.last_audio_bytes + previous_audio_chunks = conn.last_audio_chunks + turn_index = conn.turns_completed + 1 + connection_age_ms = conn.age_ms + gap_timeout_ms = round(config.gap_timeout_s * 1000) + + while True: + receive_timeout_s = config.request_timeout_s + if audio_chunks > 0 and config.gap_timeout_s > 0: + receive_timeout_s = min(receive_timeout_s, config.gap_timeout_s) + + try: + msg = await asyncio.wait_for(conn.ws.receive(), timeout=receive_timeout_s) + except asyncio.TimeoutError: + if audio_chunks > 0 and config.gap_timeout_s > 0: + gap_ms = max(0, round((time.perf_counter() - last_audio_at) * 1000)) + max_gap_ms = max(max_gap_ms, gap_ms) + outcome = "gap_timeout" + error = f"no audio.delta for {gap_ms}ms" + else: + outcome = "no_audio_timeout" + error = f"no audio.delta before {round(config.request_timeout_s * 1000)}ms" + break + + now = time.perf_counter() + if msg.type in ( + aiohttp.WSMsgType.CLOSED, + aiohttp.WSMsgType.CLOSE, + aiohttp.WSMsgType.CLOSING, + aiohttp.WSMsgType.ERROR, + ): + outcome = "ws_closed_after_audio" if audio_chunks else "ws_closed_no_audio" + error = f"websocket message type {msg.type.name}" + break + + if msg.type != aiohttp.WSMsgType.TEXT: + continue + + try: + data = json.loads(str(msg.data)) + except json.JSONDecodeError as exc: + outcome = "invalid_json" + error = str(exc) + break + + msg_type = data.get("type") + if msg_type == "audio.delta": + try: + chunk = base64.b64decode(str(data.get("delta") or "")) + except (binascii.Error, ValueError) as exc: + outcome = "invalid_audio_delta" + error = str(exc) + break + + if not chunk: + empty_audio_deltas += 1 + continue + + if audio_chunks == 0: + first_audio_ms = max(0, round((now - turn_started_at) * 1000)) + request_to_first_audio_ms = max(0, round((now - send_started_at) * 1000)) + total_to_first_audio_ms = max(0, round((now - job_started_at) * 1000)) + else: + gap_ms = max(0, round((now - last_audio_at) * 1000)) + max_gap_ms = max(max_gap_ms, gap_ms) + + audio_chunks += 1 + audio_bytes += len(chunk) + last_audio_at = now + if output_audio_dir is not None: + audio.extend(chunk) + elif msg_type == "audio.done": + trace_id = str(data.get("trace_id") or "") + outcome = "success" if audio_chunks else "no_audio_done" + break + elif msg_type == "audio.clear": + outcome = "audio_clear" + break + elif msg_type == "error": + trace_id = str(data.get("trace_id") or "") + outcome = "api_error" + error = str(data.get("message") or data) + break + + elapsed_ms = _elapsed_ms(turn_started_at) + audio_path = "" + if output_audio_dir is not None and audio: + audio_path = str(_write_wav(output_audio_dir / f"{iteration:04d}_{request_id}.wav", bytes(audio))) + + result = Result( + run_id=run_id, + request_id=request_id, + worker_id=worker_id, + iteration=iteration, + started_at=_now_utc(), + connection_id=conn.connection_id, + connection_age_ms=connection_age_ms, + connection_turn_index=turn_index, + reused_connection=(not opened_for_request), + opened_for_request=opened_for_request, + idle_before_turn_ms=idle_before_turn_ms, + previous_outcome=previous_outcome, + previous_audio_bytes=previous_audio_bytes, + previous_audio_chunks=previous_audio_chunks, + text_len=len(text), + connect_ms=conn.opened_connect_ms if opened_for_request else 0, + send_ms=send_ms, + first_audio_ms=first_audio_ms, + request_to_first_audio_ms=request_to_first_audio_ms, + total_to_first_audio_ms=total_to_first_audio_ms, + elapsed_ms=elapsed_ms, + audio_chunks=audio_chunks, + audio_bytes=audio_bytes, + empty_audio_deltas=empty_audio_deltas, + max_gap_ms=max_gap_ms, + gap_timeout_ms=gap_timeout_ms, + outcome=outcome, + trace_id=trace_id, + error=error, + audio_path=audio_path, + ) + conn.turns_completed += 1 + conn.last_outcome = outcome + conn.last_audio_bytes = audio_bytes + conn.last_audio_chunks = audio_chunks + return result + + +def _compact_result_line(result: Result) -> str: + return ( + "[xai-tts-load] " + f"i={result.iteration} worker={result.worker_id} conn={result.connection_id} " + f"turn={result.connection_turn_index} outcome={result.outcome} " + f"age_ms={result.connection_age_ms} idle_ms={result.idle_before_turn_ms} " + f"connect_ms={result.connect_ms} first_audio_ms={result.first_audio_ms} " + f"request_to_first_audio_ms={result.request_to_first_audio_ms} " + f"total_to_first_audio_ms={result.total_to_first_audio_ms} " + f"max_gap_ms={result.max_gap_ms} chunks={result.audio_chunks} bytes={result.audio_bytes} " + f"trace_id={result.trace_id}" + ) + + +async def _run_worker( + *, + config: Config, + run_id: str, + worker_id: int, + iterations: list[int], + jsonl_handle: Any, + write_lock: asyncio.Lock, +) -> list[Result]: + results: list[Result] = [] + conn: XAIConnection | None = None + async with aiohttp.ClientSession() as session: + for iteration in iterations: + text = config.texts[(iteration - 1) % len(config.texts)] + job_started_at = time.perf_counter() + opened_for_request = False + idle_before_turn_ms = 0 + if conn is None or conn.closed: + conn = await _open_connection(session, config, worker_id=worker_id) + opened_for_request = True + if config.prewarm_idle_s > 0: + idle_started_at = time.perf_counter() + await asyncio.sleep(config.prewarm_idle_s) + idle_before_turn_ms = _elapsed_ms(idle_started_at) + elif config.idle_between_requests_s > 0: + idle_started_at = time.perf_counter() + await asyncio.sleep(config.idle_between_requests_s) + idle_before_turn_ms = _elapsed_ms(idle_started_at) + + try: + result = await _receive_xai_turn( + config=config, + conn=conn, + text=text, + run_id=run_id, + worker_id=worker_id, + iteration=iteration, + job_started_at=job_started_at, + opened_for_request=opened_for_request, + idle_before_turn_ms=idle_before_turn_ms, + output_audio_dir=config.audio_dir, + ) + except Exception as exc: + result = Result( + run_id=run_id, + request_id=uuid.uuid4().hex[:12], + worker_id=worker_id, + iteration=iteration, + started_at=_now_utc(), + connection_id=conn.connection_id if conn is not None else "", + connection_age_ms=conn.age_ms if conn is not None else 0, + connection_turn_index=(conn.turns_completed + 1) if conn is not None else 0, + reused_connection=not opened_for_request, + opened_for_request=opened_for_request, + idle_before_turn_ms=idle_before_turn_ms, + previous_outcome=conn.last_outcome if conn is not None else "", + previous_audio_bytes=conn.last_audio_bytes if conn is not None else 0, + previous_audio_chunks=conn.last_audio_chunks if conn is not None else 0, + text_len=len(text), + connect_ms=conn.opened_connect_ms if conn is not None and opened_for_request else 0, + send_ms=0, + first_audio_ms=None, + request_to_first_audio_ms=None, + total_to_first_audio_ms=None, + elapsed_ms=_elapsed_ms(job_started_at), + audio_chunks=0, + audio_bytes=0, + empty_audio_deltas=0, + max_gap_ms=0, + gap_timeout_ms=round(config.gap_timeout_s * 1000), + outcome="exception", + trace_id="", + error=f"{type(exc).__name__}: {exc}", + audio_path="", + ) + + results.append(result) + async with write_lock: + jsonl_handle.write(json.dumps(asdict(result), ensure_ascii=False) + "\n") + jsonl_handle.flush() + print(_compact_result_line(result), flush=True) + + should_close = not config.reuse_connection + if result.outcome != "success" and config.reset_on_failure: + should_close = True + if should_close: + await _close_connection(conn, reason=result.outcome) + conn = None + + await _close_connection(conn, reason="worker_done") + return results + + +def _partition_iterations(iterations: int, concurrency: int) -> list[list[int]]: + buckets: list[list[int]] = [[] for _ in range(concurrency)] + for iteration in range(1, iterations + 1): + buckets[(iteration - 1) % concurrency].append(iteration) + return buckets + + +def _write_summary_csv(path: Path, results: list[Result]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + fields = list(asdict(results[0]).keys()) if results else list(Result.__dataclass_fields__.keys()) + with path.open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + for result in sorted(results, key=lambda item: item.iteration): + writer.writerow(asdict(result)) + + +def _metric_values(results: list[Result], field: str) -> list[int]: + values: list[int] = [] + for result in results: + value = getattr(result, field) + if isinstance(value, int): + values.append(value) + return values + + +def _print_summary(config: Config, results: list[Result], started_at: float) -> None: + outcomes = Counter(result.outcome for result in results) + first_audio_values = _metric_values(results, "first_audio_ms") + request_to_first_values = _metric_values(results, "request_to_first_audio_ms") + total_to_first_values = _metric_values(results, "total_to_first_audio_ms") + connect_values = [result.connect_ms for result in results if result.opened_for_request] + connection_age_values = _metric_values(results, "connection_age_ms") + gap_values = _metric_values(results, "max_gap_ms") + no_audio = sum(1 for result in results if result.audio_bytes == 0) + gap_timeouts = sum(1 for result in results if result.outcome == "gap_timeout") + total_near_3s = sum( + 1 + for result in results + if result.total_to_first_audio_ms is not None and result.total_to_first_audio_ms >= 2800 + ) + request_near_3s = sum( + 1 + for result in results + if result.request_to_first_audio_ms is not None and result.request_to_first_audio_ms >= 2800 + ) + + print("", flush=True) + print("[xai-tts-load] resumo", flush=True) + print(f" total={len(results)} duration_ms={_elapsed_ms(started_at)}", flush=True) + print( + f" reuse_connection={int(config.reuse_connection)} concurrency={config.concurrency} " + f"prewarm_idle_ms={round(config.prewarm_idle_s * 1000)} " + f"idle_between_requests_ms={round(config.idle_between_requests_s * 1000)}", + flush=True, + ) + print(f" outcomes={dict(outcomes)}", flush=True) + print( + f" no_audio={no_audio} gap_timeouts={gap_timeouts} " + f"request_to_first_audio_ge_2800ms={request_near_3s} " + f"total_to_first_audio_ge_2800ms={total_near_3s}", + flush=True, + ) + print( + " connection_age_ms " + f"p50={_percentile(connection_age_values, 0.50)} p95={_percentile(connection_age_values, 0.95)} " + f"max={max(connection_age_values) if connection_age_values else None}", + flush=True, + ) + print( + " connect_ms " + f"p50={_percentile(connect_values, 0.50)} p95={_percentile(connect_values, 0.95)} " + f"max={max(connect_values) if connect_values else None}", + flush=True, + ) + print( + " first_audio_ms " + f"p50={_percentile(first_audio_values, 0.50)} p95={_percentile(first_audio_values, 0.95)} " + f"max={max(first_audio_values) if first_audio_values else None}", + flush=True, + ) + print( + " request_to_first_audio_ms " + f"p50={_percentile(request_to_first_values, 0.50)} p95={_percentile(request_to_first_values, 0.95)} " + f"max={max(request_to_first_values) if request_to_first_values else None}", + flush=True, + ) + print( + " total_to_first_audio_ms " + f"p50={_percentile(total_to_first_values, 0.50)} p95={_percentile(total_to_first_values, 0.95)} " + f"max={max(total_to_first_values) if total_to_first_values else None}", + flush=True, + ) + print( + " max_gap_ms " + f"p50={_percentile(gap_values, 0.50)} p95={_percentile(gap_values, 0.95)} " + f"max={max(gap_values) if gap_values else None}", + flush=True, + ) + + +def _load_texts(args: argparse.Namespace) -> tuple[str, ...]: + texts = [item for item in (args.text or []) if item] + if args.texts_file: + for line in Path(args.texts_file).read_text(encoding="utf-8").splitlines(): + stripped = line.strip() + if stripped and not stripped.startswith("#"): + texts.append(stripped) + if not texts: + texts.append(DEFAULT_TEXT) + return tuple(texts) + + +def parse_args(argv: list[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Carga direta no WebSocket de TTS xAI, sem LiveKit/TIA.", + ) + parser.add_argument("--env-file", default=".env.dev") + parser.add_argument("--websocket-url", default="") + parser.add_argument("--api-key", default="") + parser.add_argument("--voice", default="") + parser.add_argument("--language", default="") + parser.add_argument("--iterations", type=int, default=20) + parser.add_argument("--concurrency", type=int, default=1) + mode = parser.add_mutually_exclusive_group() + mode.add_argument("--reuse-connection", dest="reuse_connection", action="store_true", default=True) + mode.add_argument("--new-connection-per-request", dest="reuse_connection", action="store_false") + parser.add_argument("--keep-connection-after-failure", action="store_true") + parser.add_argument( + "--prewarm-idle-ms", + type=int, + default=0, + help="Abre o WebSocket e espera este tempo antes do primeiro texto da conexao.", + ) + parser.add_argument( + "--idle-between-requests-ms", + type=int, + default=0, + help="Espera este tempo entre turnos quando a mesma conexao e reutilizada.", + ) + parser.add_argument("--request-timeout-ms", type=int, default=10_000) + parser.add_argument("--connect-timeout-ms", type=int, default=10_000) + parser.add_argument("--gap-timeout-ms", type=int, default=3_000) + parser.add_argument("--optimize-streaming-latency", type=int, default=1) + parser.add_argument("--no-text-normalization", action="store_true") + parser.add_argument("--text", action="append", help="Texto a sintetizar. Pode repetir.") + parser.add_argument("--texts-file", default="", help="Arquivo com um texto por linha.") + parser.add_argument("--output-jsonl", default=str(DEFAULT_OUTPUT_JSONL)) + parser.add_argument("--summary-csv", default=str(DEFAULT_SUMMARY_CSV)) + parser.add_argument("--no-summary-csv", action="store_true") + parser.add_argument("--write-audio-dir", default="") + return parser.parse_args(argv) + + +def load_config(args: argparse.Namespace) -> Config: + env_file = Path(args.env_file) + if env_file.exists(): + load_dotenv(env_file, override=False) + + api_key = str(args.api_key or _env_str("XAI_API_KEY")).strip() + websocket_url = str(args.websocket_url or _env_str("XAI_WEBSOCKET_URL")).strip() + if not api_key: + raise RuntimeError("Missing XAI_API_KEY. Defina no .env.dev ou passe --api-key.") + if not websocket_url: + raise RuntimeError("Missing XAI_WEBSOCKET_URL. Defina no .env.dev ou passe --websocket-url.") + + return Config( + env_file=env_file, + api_key=api_key, + websocket_url=websocket_url, + voice=str(args.voice or _env_str("XAI_TTS_VOICE", "ara")).strip() or "ara", + language=str(args.language or _env_str("XAI_TTS_LANGUAGE", "pt-BR")).strip() or "pt-BR", + text_normalization=not bool(args.no_text_normalization), + optimize_streaming_latency=max(0, int(args.optimize_streaming_latency)), + iterations=max(1, int(args.iterations)), + concurrency=max(1, int(args.concurrency)), + reuse_connection=bool(args.reuse_connection), + reset_on_failure=not bool(args.keep_connection_after_failure), + prewarm_idle_s=max(0.0, int(args.prewarm_idle_ms) / 1000), + idle_between_requests_s=max(0.0, int(args.idle_between_requests_ms) / 1000), + request_timeout_s=max(0.1, int(args.request_timeout_ms) / 1000), + connect_timeout_s=max(0.1, int(args.connect_timeout_ms) / 1000), + gap_timeout_s=max(0.0, int(args.gap_timeout_ms) / 1000), + output_jsonl=Path(args.output_jsonl), + summary_csv=None if args.no_summary_csv else Path(args.summary_csv), + audio_dir=Path(args.write_audio_dir) if args.write_audio_dir else None, + texts=_load_texts(args), + ) + + +async def run(argv: list[str] | None = None) -> int: + config = load_config(parse_args(argv)) + config.output_jsonl.parent.mkdir(parents=True, exist_ok=True) + if config.audio_dir is not None: + config.audio_dir.mkdir(parents=True, exist_ok=True) + + run_id = uuid.uuid4().hex[:12] + started_at = time.perf_counter() + print( + "[xai-tts-load] iniciando " + f"run_id={run_id} iterations={config.iterations} concurrency={config.concurrency} " + f"reuse_connection={int(config.reuse_connection)} voice={config.voice} " + f"language={config.language} prewarm_idle_ms={round(config.prewarm_idle_s * 1000)} " + f"idle_between_requests_ms={round(config.idle_between_requests_s * 1000)} " + f"output={config.output_jsonl}", + flush=True, + ) + + write_lock = asyncio.Lock() + partitions = _partition_iterations(config.iterations, config.concurrency) + with config.output_jsonl.open("w", encoding="utf-8") as jsonl_handle: + tasks = [ + _run_worker( + config=config, + run_id=run_id, + worker_id=index + 1, + iterations=iterations, + jsonl_handle=jsonl_handle, + write_lock=write_lock, + ) + for index, iterations in enumerate(partitions) + if iterations + ] + nested_results = await asyncio.gather(*tasks) + + results = [result for batch in nested_results for result in batch] + results.sort(key=lambda item: item.iteration) + if config.summary_csv is not None: + _write_summary_csv(config.summary_csv, results) + print(f"[xai-tts-load] csv={config.summary_csv}", flush=True) + _print_summary(config, results, started_at) + return 1 if any(result.outcome != "success" for result in results) else 0 + + +def main() -> None: + raise SystemExit(asyncio.run(run())) + + +if __name__ == "__main__": + main() diff --git a/src/app/tools/xai_tts_ws.py b/src/app/tools/xai_tts_ws.py new file mode 100644 index 0000000..3e0f92a --- /dev/null +++ b/src/app/tools/xai_tts_ws.py @@ -0,0 +1,676 @@ +"""Captura crua dos eventos do WebSocket xAI TTS em JSONL. + +Este coletor nao usa LiveKit, AudioEmitter, buffer, retry nem timeout de gap. +Cada linha do JSONL e um evento do transporte/provedor, identificado por call_id. +""" + +from __future__ import annotations + +import argparse +import asyncio +import base64 +import binascii +from collections import Counter +import hashlib +import json +import os +import time +import uuid +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from urllib.parse import urlencode + +import aiohttp +from dotenv import load_dotenv + +SAMPLE_RATE = 24_000 +PCM_BYTES_PER_SAMPLE = 2 +DEFAULT_TEXT = "Ola, esta e uma captura bruta de eventos do websocket de TTS." + + +@dataclass(frozen=True, slots=True) +class Config: + api_key: str + websocket_url: str + voice: str + language: str + text: str + calls: int + concurrency: int + prewarm: bool + event_timeout_s: float | None + output: Path + summary_output: Path + include_delta_base64: bool + initial_playout_buffer_ms: float + + +def utc_now() -> str: + return datetime.now(timezone.utc).isoformat(timespec="milliseconds") + + +def elapsed_ms(started_at: float) -> int: + return max(0, round((time.perf_counter() - started_at) * 1000)) + + +def bool_query(value: bool) -> str: + return "true" if value else "false" + + +def build_url(config: Config) -> str: + query = { + "voice": config.voice, + "language": config.language, + "codec": "pcm", + "sample_rate": str(SAMPLE_RATE), + "optimize_streaming_latency": "1", + "text_normalization": bool_query(True), + } + return f"{config.websocket_url}?{urlencode(query)}" + + +def parser() -> argparse.ArgumentParser: + result = argparse.ArgumentParser( + description="Captura eventos brutos do WebSocket xAI TTS em JSONL, sem LiveKit." + ) + result.add_argument("--env-file", default=".env.dev") + result.add_argument("--api-key", default="") + result.add_argument("--websocket-url", default="") + result.add_argument("--voice", default="") + result.add_argument("--language", default="") + result.add_argument("--text", default=DEFAULT_TEXT) + result.add_argument("--calls", type=int, default=1) + result.add_argument("--concurrency", type=int, default=1) + result.add_argument( + "--prewarm", + action="store_true", + help="Abre todas as conexoes antes da janela medida. Requer --calls igual a --concurrency.", + ) + result.add_argument( + "--event-timeout-ms", + type=int, + default=65_000, + help="Limite de seguranca uniforme para cada ws.receive; 0 desativa. Nao e timeout de gap.", + ) + result.add_argument("--output", default=".run/xai-tts-ws-capture/events.jsonl") + result.add_argument( + "--summary-output", + default="", + help="JSON summary output. Default: .summary.json.", + ) + result.add_argument( + "--initial-playout-buffer-ms", + type=float, + default=200.0, + help="Initial buffer used only for the underflow simulation in the summary (default: 200).", + ) + result.add_argument( + "--include-delta-base64", + action="store_true", + help="Inclui o campo delta integral em cada audio.delta; pode gerar arquivos muito grandes.", + ) + return result + + +def load_config(args: argparse.Namespace) -> Config: + env_file = Path(args.env_file) + if env_file.exists(): + load_dotenv(env_file, override=False) + api_key = str(args.api_key or os.getenv("XAI_API_KEY", "")).strip() + websocket_url = str(args.websocket_url or os.getenv("XAI_WEBSOCKET_URL", "")).strip() + if not api_key: + raise RuntimeError("XAI_API_KEY ausente; informe --api-key ou --env-file.") + if not websocket_url: + raise RuntimeError("XAI_WEBSOCKET_URL ausente; informe --websocket-url ou --env-file.") + timeout_ms = max(0, int(args.event_timeout_ms)) + output = Path(args.output) + summary_output = Path(args.summary_output) if args.summary_output else output.with_suffix(".summary.json") + return Config( + api_key=api_key, + websocket_url=websocket_url, + voice=str(args.voice or os.getenv("XAI_TTS_VOICE", "ara")).strip() or "ara", + language=str(args.language or os.getenv("XAI_TTS_LANGUAGE", "pt-BR")).strip() or "pt-BR", + text=str(args.text), + calls=max(1, int(args.calls)), + concurrency=max(1, int(args.concurrency)), + prewarm=bool(args.prewarm), + event_timeout_s=(timeout_ms / 1000) if timeout_ms else None, + output=output, + summary_output=summary_output, + include_delta_base64=bool(args.include_delta_base64), + initial_playout_buffer_ms=max(0.0, float(args.initial_playout_buffer_ms)), + ) + + +class JsonlWriter: + def __init__(self, output: Path) -> None: + output.parent.mkdir(parents=True, exist_ok=True) + self._handle = output.open("w", encoding="utf-8") + self._lock = asyncio.Lock() + + async def write(self, record: dict[str, Any]) -> None: + async with self._lock: + self._handle.write(json.dumps(record, ensure_ascii=False, separators=(",", ":")) + "\n") + self._handle.flush() + + def close(self) -> None: + self._handle.close() + + +async def receive_message(ws: aiohttp.ClientWebSocketResponse, timeout_s: float | None) -> aiohttp.WSMessage: + if timeout_s is None: + return await ws.receive() + return await asyncio.wait_for(ws.receive(), timeout=timeout_s) + + +def event_base(*, run_id: str, call_id: str, call_number: int, sequence: int, started_at: float) -> dict[str, Any]: + return { + "run_id": run_id, + "call_id": call_id, + "call_number": call_number, + "sequence": sequence, + "received_at_utc": utc_now(), + "received_after_call_start_ms": elapsed_ms(started_at), + } + + +async def capture_call( + *, + session: aiohttp.ClientSession, + config: Config, + writer: JsonlWriter, + run_id: str, + call_number: int, + call_id: str, + limiter: asyncio.Semaphore, + prewarmed_ws: aiohttp.ClientWebSocketResponse | None = None, +) -> str: + async with limiter: + started_at = time.perf_counter() + sequence = 0 + status = "unknown" + ws = prewarmed_ws + + async def write(kind: str, **fields: Any) -> None: + nonlocal sequence + sequence += 1 + record = event_base( + run_id=run_id, + call_id=call_id, + call_number=call_number, + sequence=sequence, + started_at=started_at, + ) + record["event"] = kind + record.update(fields) + await writer.write(record) + + try: + if ws is None: + await write("client.connecting", websocket_url=config.websocket_url) + connect_started_at = time.perf_counter() + ws = await session.ws_connect( + build_url(config), + headers={"Authorization": f"Bearer {config.api_key}"}, + heartbeat=20, + ) + await write("client.connected", connect_ms=elapsed_ms(connect_started_at)) + else: + await write("client.using_prewarmed_connection") + + await ws.send_json({"type": "text.clear"}) + await write("client.sent", message_type="text.clear") + clear_acknowledged = False + while not clear_acknowledged: + message = await receive_message(ws, config.event_timeout_s) + now = time.perf_counter() + if message.type != aiohttp.WSMsgType.TEXT: + await write("transport.message", ws_type=message.type.name, data=str(message.data)) + raise RuntimeError(f"esperava audio.clear, recebeu WebSocket {message.type.name}") + try: + payload = json.loads(str(message.data)) + except json.JSONDecodeError: + await write("provider.invalid_json", raw_text=str(message.data)) + raise + await write( + "provider.event", + message_type=str(payload.get("type") or ""), + provider_message=payload, + received_monotonic_ms=round(now * 1000), + ) + clear_acknowledged = payload.get("type") == "audio.clear" + if payload.get("type") == "error": + raise RuntimeError(str(payload.get("message") or payload)) + + await ws.send_json({"type": "text.delta", "delta": config.text}) + await write("client.sent", message_type="text.delta", text_chars=len(config.text)) + await ws.send_json({"type": "text.done"}) + await write("client.sent", message_type="text.done") + + while True: + try: + message = await receive_message(ws, config.event_timeout_s) + except asyncio.TimeoutError: + await write("client.receive_timeout", timeout_ms=None if config.event_timeout_s is None else round(config.event_timeout_s * 1000)) + status = "receive_timeout" + break + + now = time.perf_counter() + if message.type != aiohttp.WSMsgType.TEXT: + await write("transport.message", ws_type=message.type.name, data=str(message.data)) + status = f"transport_{message.type.name.lower()}" + break + + try: + payload = json.loads(str(message.data)) + except json.JSONDecodeError: + await write("provider.invalid_json", raw_text=str(message.data)) + status = "invalid_json" + break + + message_type = str(payload.get("type") or "") + if message_type == "audio.delta": + encoded = str(payload.get("delta") or "") + try: + pcm = base64.b64decode(encoded, validate=True) + except (binascii.Error, ValueError) as exc: + await write( + "provider.audio_delta_invalid", + message_type=message_type, + base64_characters=len(encoded), + error=str(exc), + provider_message={key: value for key, value in payload.items() if key != "delta"}, + ) + status = "invalid_audio_delta" + break + record: dict[str, Any] = { + "message_type": message_type, + "base64_characters": len(encoded), + "audio_bytes": len(pcm), + "audio_ms": round(len(pcm) * 1000 / (SAMPLE_RATE * PCM_BYTES_PER_SAMPLE), 3), + "empty_audio": not bool(pcm), + "audio_sha256": hashlib.sha256(pcm).hexdigest(), + "provider_message": {key: value for key, value in payload.items() if key != "delta"}, + "received_monotonic_ms": round(now * 1000), + } + if config.include_delta_base64: + record["delta_base64"] = encoded + await write("provider.audio_delta", **record) + continue + + await write( + "provider.event", + message_type=message_type, + provider_message=payload, + received_monotonic_ms=round(now * 1000), + ) + if message_type == "audio.done": + status = "audio_done" + break + if message_type == "error": + status = "provider_error" + break + except Exception as exc: + if status == "unknown": + status = "exception" + await write("client.exception", error_type=type(exc).__name__, error=str(exc)) + finally: + if ws is not None and not ws.closed: + await ws.close() + await write("client.closed") + await write("client.finished", status=status, elapsed_ms=elapsed_ms(started_at)) + return status + + +async def prewarm_connection( + *, + session: aiohttp.ClientSession, + config: Config, + writer: JsonlWriter, + run_id: str, + call_number: int, + call_id: str, +) -> aiohttp.ClientWebSocketResponse: + started_at = time.perf_counter() + await writer.write({ + **event_base(run_id=run_id, call_id=call_id, call_number=call_number, sequence=-1, started_at=started_at), + "event": "client.prewarm_connecting", + "websocket_url": config.websocket_url, + }) + try: + ws = await session.ws_connect( + build_url(config), + headers={"Authorization": f"Bearer {config.api_key}"}, + heartbeat=20, + ) + except Exception as exc: + await writer.write({ + **event_base(run_id=run_id, call_id=call_id, call_number=call_number, sequence=0, started_at=started_at), + "event": "client.prewarm_failed", + "error_type": type(exc).__name__, + "error": str(exc), + }) + raise + await writer.write({ + **event_base(run_id=run_id, call_id=call_id, call_number=call_number, sequence=0, started_at=started_at), + "event": "client.prewarm_ready", + "connect_ms": elapsed_ms(started_at), + }) + return ws + + +def percentile(values: list[float], percentile_value: float) -> float | None: + """Return the nearest-rank percentile without introducing a dependency.""" + if not values: + return None + ordered = sorted(values) + index = round((len(ordered) - 1) * percentile_value) + return ordered[max(0, min(index, len(ordered) - 1))] + + +def round_ms(value: float | None) -> float | None: + return None if value is None else round(value, 3) + + +def replay_underflow( + audio_frames: list[tuple[float, float]], initial_buffer_ms: float +) -> tuple[int, float, float]: + """Replay raw arrivals against real-time playout after the initial buffer fills.""" + buffered_ms = 0.0 + playout_started = False + previous_arrival_ms: float | None = None + events = 0 + total_underflow_ms = 0.0 + largest_underflow_ms = 0.0 + + for arrived_at_ms, audio_ms in audio_frames: + if playout_started: + assert previous_arrival_ms is not None + elapsed_ms = arrived_at_ms - previous_arrival_ms + missing_audio_ms = max(0.0, elapsed_ms - buffered_ms) + if missing_audio_ms: + events += 1 + total_underflow_ms += missing_audio_ms + largest_underflow_ms = max(largest_underflow_ms, missing_audio_ms) + buffered_ms = max(0.0, buffered_ms - elapsed_ms) + audio_ms + previous_arrival_ms = arrived_at_ms + continue + + buffered_ms += audio_ms + if buffered_ms >= initial_buffer_ms: + playout_started = True + previous_arrival_ms = arrived_at_ms + + return events, total_underflow_ms, largest_underflow_ms + + +def metric_distribution(values: list[float]) -> dict[str, float | int | None]: + return { + "count": len(values), + "p50": round_ms(percentile(values, 0.50)), + "p95": round_ms(percentile(values, 0.95)), + "max": round_ms(max(values)) if values else None, + } + + +def summarize_capture(config: Config, run_id: str) -> dict[str, Any]: + """Build a provider-facing summary exclusively from the raw JSONL capture.""" + calls: dict[str, list[dict[str, Any]]] = {} + run_events: list[dict[str, Any]] = [] + with config.output.open(encoding="utf-8") as handle: + for line in handle: + record = json.loads(line) + call_id = record.get("call_id") + if call_id: + calls.setdefault(str(call_id), []).append(record) + else: + run_events.append(record) + + status_counts: Counter[str] = Counter() + error_counts: Counter[str] = Counter() + ttfb_ms: list[float] = [] + rtf: list[float] = [] + audio_duration_ms: list[float] = [] + frame_duration_ms: list[float] = [] + frame_gaps_ms: list[float] = [] + per_call_max_gaps_ms: list[float] = [] + frames_per_call: list[float] = [] + calls_with_audio = 0 + calls_with_underflow = 0 + underflow_events = 0 + total_underflow_ms = 0.0 + largest_underflow_ms = 0.0 + empty_audio_deltas = 0 + + for records in calls.values(): + records.sort(key=lambda record: int(record.get("sequence", -2))) + for record in records: + event = str(record.get("event") or "") + message_type = str(record.get("message_type") or "") + if event == "client.finished": + status_counts[str(record.get("status") or "unknown")] += 1 + if event == "client.prewarm_failed": + error_counts["prewarm_failed"] += 1 + elif event == "client.exception": + error_counts["client_exception"] += 1 + elif event == "client.receive_timeout": + error_counts["receive_timeout"] += 1 + elif event == "provider.audio_delta_invalid": + error_counts["invalid_audio_delta"] += 1 + elif event == "provider.invalid_json": + error_counts["invalid_json"] += 1 + elif event == "transport.message": + error_counts[f"transport_{str(record.get('ws_type') or 'unknown').lower()}"] += 1 + elif event == "provider.event" and message_type == "error": + error_counts["provider_error"] += 1 + elif event == "provider.audio_delta" and record.get("empty_audio"): + empty_audio_deltas += 1 + + text_done_at_ms = next( + ( + float(record["received_after_call_start_ms"]) + for record in records + if record.get("event") == "client.sent" and record.get("message_type") == "text.done" + ), + None, + ) + audio_frames = [ + (float(record["received_after_call_start_ms"]), float(record["audio_ms"])) + for record in records + if record.get("event") == "provider.audio_delta" and float(record.get("audio_bytes") or 0) > 0 + ] + if text_done_at_ms is None or not audio_frames: + continue + + calls_with_audio += 1 + frames_per_call.append(float(len(audio_frames))) + duration_ms = sum(frame[1] for frame in audio_frames) + audio_duration_ms.append(duration_ms) + frame_duration_ms.extend(frame[1] for frame in audio_frames) + ttfb_ms.append(audio_frames[0][0] - text_done_at_ms) + if duration_ms: + rtf.append((audio_frames[-1][0] - text_done_at_ms) / duration_ms) + gaps = [audio_frames[index][0] - audio_frames[index - 1][0] for index in range(1, len(audio_frames))] + frame_gaps_ms.extend(gaps) + if gaps: + per_call_max_gaps_ms.append(max(gaps)) + events, total_ms, largest_ms = replay_underflow(audio_frames, config.initial_playout_buffer_ms) + if events: + calls_with_underflow += 1 + underflow_events += events + total_underflow_ms += total_ms + largest_underflow_ms = max(largest_underflow_ms, largest_ms) + + run_failed = next((record for record in run_events if record.get("event") == "run.failed"), None) + completed_calls = status_counts.get("audio_done", 0) + observed_errors = sum(error_counts.values()) + return { + "summary_version": 1, + "language": "en", + "source_events_jsonl": str(config.output), + "run_id": run_id, + "configuration": { + "requested_calls": config.calls, + "concurrency": config.concurrency, + "prewarm": config.prewarm, + "initial_playout_buffer_ms": config.initial_playout_buffer_ms, + "event_timeout_ms": None if config.event_timeout_s is None else round(config.event_timeout_s * 1000), + "pcm_format": f"{SAMPLE_RATE} Hz, mono, signed 16-bit PCM", + }, + "completion": { + "calls_observed": len(calls), + "calls_with_audio": calls_with_audio, + "audio_done": completed_calls, + "audio_done_rate": round(completed_calls / config.calls, 4) if config.calls else None, + "finished_statuses": dict(sorted(status_counts.items())), + "run_failed": bool(run_failed), + "run_failure": None if run_failed is None else run_failed.get("error"), + }, + "latency": { + "ttfb_ms": metric_distribution(ttfb_ms), + "rtf": metric_distribution(rtf), + "generated_audio_ms_per_call": metric_distribution(audio_duration_ms), + }, + "frame_delivery": { + "nonempty_audio_frames_per_call": metric_distribution(frames_per_call), + "audio_ms_per_frame": metric_distribution(frame_duration_ms), + "interframe_gap_ms": metric_distribution(frame_gaps_ms), + "per_call_max_interframe_gap_ms": metric_distribution(per_call_max_gaps_ms), + "empty_audio_delta_count": empty_audio_deltas, + }, + "simulated_playout": { + "definition": "Raw audio.delta arrivals replayed at real-time PCM consumption after the initial buffer is accumulated.", + "initial_buffer_ms": config.initial_playout_buffer_ms, + "calls_with_underflow": calls_with_underflow, + "calls_with_underflow_rate": round(calls_with_underflow / config.calls, 4) if config.calls else None, + "underflow_events": underflow_events, + "total_underflow_ms": round_ms(total_underflow_ms), + "largest_underflow_ms": round_ms(largest_underflow_ms), + }, + "errors": { + "total_observed": observed_errors, + "by_type": dict(sorted(error_counts.items())), + "empty_audio_delta_count": empty_audio_deltas, + }, + } + + +def format_summary(summary: dict[str, Any]) -> str: + """Create a concise, copy-pasteable English summary for terminal output.""" + configuration = summary["configuration"] + completion = summary["completion"] + latency = summary["latency"] + frames = summary["frame_delivery"] + playout = summary["simulated_playout"] + errors = summary["errors"] + ttfb = latency["ttfb_ms"] + rtf = latency["rtf"] + gaps = frames["interframe_gap_ms"] + return "\n".join( + [ + "TTS WebSocket Capture Summary", + f"Load: {configuration['requested_calls']} calls at C={configuration['concurrency']} (prewarm: {'on' if configuration['prewarm'] else 'off'})", + f"Completion: {completion['audio_done']}/{configuration['requested_calls']} audio.done ({completion['audio_done_rate']:.1%})", + f"TTFB (text.done -> first non-empty audio): p50 {ttfb['p50']} ms | p95 {ttfb['p95']} ms | max {ttfb['max']} ms", + f"RTF (last audio arrival / generated audio): p50 {rtf['p50']} | p95 {rtf['p95']} | max {rtf['max']}", + f"Inter-frame gap: p50 {gaps['p50']} ms | p95 {gaps['p95']} ms | max {gaps['max']} ms", + f"Simulated underflow ({playout['initial_buffer_ms']} ms initial buffer): {playout['calls_with_underflow']}/{configuration['requested_calls']} calls | {playout['underflow_events']} events | largest {playout['largest_underflow_ms']} ms | total {playout['total_underflow_ms']} ms", + f"Errors: {errors['total_observed']} observed | empty audio.delta frames: {errors['empty_audio_delta_count']} | by type: {errors['by_type'] or 'none'}", + ] + ) + + +async def run(argv: list[str] | None = None) -> int: + config = load_config(parser().parse_args(argv)) + run_id = uuid.uuid4().hex + writer = JsonlWriter(config.output) + limiter = asyncio.Semaphore(config.concurrency) + statuses: list[str] = [] + call_ids = {number: f"call-{number:04d}-{uuid.uuid4().hex[:8]}" for number in range(1, config.calls + 1)} + try: + if config.prewarm and config.calls != config.concurrency: + raise ValueError("--prewarm requer --calls igual a --concurrency para manter a mesma janela de carga.") + await writer.write( + { + "event": "run.started", + "run_id": run_id, + "started_at_utc": utc_now(), + "calls": config.calls, + "concurrency": config.concurrency, + "prewarm": config.prewarm, + "event_timeout_ms": None if config.event_timeout_s is None else round(config.event_timeout_s * 1000), + "include_delta_base64": config.include_delta_base64, + "pcm_format": {"sample_rate": SAMPLE_RATE, "channels": 1, "sample_width_bytes": PCM_BYTES_PER_SAMPLE}, + } + ) + timeout = aiohttp.ClientTimeout(total=None) + connector = aiohttp.TCPConnector(limit=config.concurrency) + async with aiohttp.ClientSession(timeout=timeout, connector=connector) as session: + prewarmed: dict[int, aiohttp.ClientWebSocketResponse | None] = {number: None for number in call_ids} + if config.prewarm: + prewarm_started_at = time.perf_counter() + results = await asyncio.gather( + *[ + prewarm_connection(session=session, config=config, writer=writer, run_id=run_id, call_number=number, call_id=call_ids[number]) + for number in call_ids + ], + return_exceptions=True, + ) + failures = [item for item in results if isinstance(item, BaseException)] + if failures: + for item in results: + if isinstance(item, aiohttp.ClientWebSocketResponse) and not item.closed: + await item.close() + raise RuntimeError(f"prewarm falhou em {len(failures)}/{config.calls} conexoes") + prewarmed = {number: results[number - 1] for number in call_ids} + await writer.write({ + "event": "run.prewarm_finished", "run_id": run_id, "finished_at_utc": utc_now(), + "prewarm_duration_ms": elapsed_ms(prewarm_started_at), "connections_ready": len(prewarmed), + }) + statuses = await asyncio.gather( + *[ + capture_call( + session=session, config=config, writer=writer, run_id=run_id, + call_number=call_number, call_id=call_ids[call_number], limiter=limiter, + prewarmed_ws=prewarmed[call_number], + ) + for call_number in call_ids + ] + ) + await writer.write( + { + "event": "run.finished", + "run_id": run_id, + "finished_at_utc": utc_now(), + "statuses": {status: statuses.count(status) for status in sorted(set(statuses))}, + } + ) + except BaseException as exc: + await writer.write( + { + "event": "run.failed", + "run_id": run_id, + "failed_at_utc": utc_now(), + "error_type": type(exc).__name__, + "error": str(exc), + } + ) + raise + finally: + writer.close() + summary = summarize_capture(config, run_id) + config.summary_output.parent.mkdir(parents=True, exist_ok=True) + config.summary_output.write_text(json.dumps(summary, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") + print(format_summary(summary)) + print(f"Summary JSON: {config.summary_output}") + return 0 if statuses and all(status == "audio_done" for status in statuses) else 1 + + +def main() -> None: + raise SystemExit(asyncio.run(run())) + + +if __name__ == "__main__": + main() diff --git a/src/app/tools/xai_tts_ws_capture.py b/src/app/tools/xai_tts_ws_capture.py new file mode 100644 index 0000000..f78751a --- /dev/null +++ b/src/app/tools/xai_tts_ws_capture.py @@ -0,0 +1,949 @@ +"""Capture raw xAI TTS WebSocket events as JSONL. + +This collector records direct client/provider traffic. It does not use an audio +emitter, retries, or a gap timeout; each JSONL line is an event for one call. +""" + +from __future__ import annotations + +import argparse +import asyncio +import base64 +import binascii +import hashlib +import json +import math +import os +import sys +import time +import uuid +from collections import Counter +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from urllib.parse import urlencode + +import aiohttp +from dotenv import load_dotenv + +SAMPLE_RATE = 24_000 +PCM_BYTES_PER_SAMPLE = 2 +DEFAULT_TEXT = "This is a raw capture of TTS WebSocket events." + + +@dataclass(frozen=True, slots=True) +class Config: + api_key: str + websocket_url: str + voice: str + language: str + text: str + calls: int + concurrency: int + prewarm: bool + event_timeout_s: float | None + connect_timeout_s: float | None + output: Path + summary_output: Path + include_delta_base64: bool + initial_playout_buffer_ms: float + + +def utc_now() -> str: + return datetime.now(timezone.utc).isoformat(timespec="milliseconds") + + +def elapsed_ms(started_at: float) -> int: + return max(0, round((time.perf_counter() - started_at) * 1000)) + + +def bool_query(value: bool) -> str: + return "true" if value else "false" + + +def build_url(config: Config) -> str: + query = { + "voice": config.voice, + "language": config.language, + "codec": "pcm", + "sample_rate": str(SAMPLE_RATE), + "optimize_streaming_latency": "1", + "text_normalization": bool_query(True), + } + return f"{config.websocket_url}?{urlencode(query)}" + + +def parser() -> argparse.ArgumentParser: + result = argparse.ArgumentParser( + description="Capture raw xAI TTS WebSocket events as JSONL." + ) + result.add_argument("--env-file", default=".env.dev") + result.add_argument("--api-key", default="") + result.add_argument("--websocket-url", default="") + result.add_argument("--voice", default="") + result.add_argument("--language", default="") + result.add_argument("--text", default=DEFAULT_TEXT) + result.add_argument("--calls", type=int, default=1) + result.add_argument("--concurrency", type=int, default=1) + result.add_argument( + "--prewarm", + action="store_true", + help="Open every WebSocket before the measured window. Requires --calls to equal --concurrency.", + ) + result.add_argument( + "--event-timeout-ms", + type=int, + default=65_000, + help="Safety limit for each ws.receive call; 0 disables it. This is not a gap timeout.", + ) + result.add_argument( + "--connect-timeout-ms", + type=int, + default=15_000, + help="Safety limit for each WebSocket handshake; 0 disables it.", + ) + result.add_argument("--output", default=".run/xai-tts-ws-capture/events.jsonl") + result.add_argument( + "--summary-output", + default="", + help="JSON summary output. Default: .summary.json.", + ) + result.add_argument( + "--initial-playout-buffer-ms", + type=float, + default=0.0, + help="Optional intentional prebuffer for the underflow simulation (default: 0; playout starts at the first non-empty audio.delta).", + ) + result.add_argument( + "--include-delta-base64", + action="store_true", + help="Include the full Base64 delta in each audio.delta event. This can produce very large files.", + ) + return result + + +def load_config(args: argparse.Namespace) -> Config: + env_file = Path(args.env_file) + if env_file.exists(): + load_dotenv(env_file, override=False) + api_key = str(args.api_key or os.getenv("XAI_API_KEY", "")).strip() + websocket_url = str( + args.websocket_url or os.getenv("XAI_WEBSOCKET_URL", "") + ).strip() + if not api_key: + raise RuntimeError("XAI_API_KEY is missing; provide --api-key or --env-file.") + if not websocket_url: + raise RuntimeError( + "XAI_WEBSOCKET_URL is missing; provide --websocket-url or --env-file." + ) + if not str(args.text).strip(): + raise ValueError("--text must not be empty.") + if int(args.calls) < 1: + raise ValueError("--calls must be at least 1.") + if int(args.concurrency) < 1: + raise ValueError("--concurrency must be at least 1.") + timeout_ms = max(0, int(args.event_timeout_ms)) + connect_timeout_ms = max(0, int(args.connect_timeout_ms)) + initial_playout_buffer_ms = float(args.initial_playout_buffer_ms) + if not math.isfinite(initial_playout_buffer_ms) or initial_playout_buffer_ms < 0: + raise ValueError( + "--initial-playout-buffer-ms must be a finite value greater than or equal to 0." + ) + output = Path(args.output) + summary_output = ( + Path(args.summary_output) + if args.summary_output + else output.with_suffix(".summary.json") + ) + if output.absolute() == summary_output.absolute(): + raise ValueError("--output and --summary-output must refer to different files.") + return Config( + api_key=api_key, + websocket_url=websocket_url, + voice=str(args.voice or os.getenv("XAI_TTS_VOICE", "ara")).strip() or "ara", + language=str(args.language or os.getenv("XAI_TTS_LANGUAGE", "pt-BR")).strip() + or "pt-BR", + text=str(args.text), + calls=int(args.calls), + concurrency=int(args.concurrency), + prewarm=bool(args.prewarm), + event_timeout_s=(timeout_ms / 1000) if timeout_ms else None, + connect_timeout_s=(connect_timeout_ms / 1000) if connect_timeout_ms else None, + output=output, + summary_output=summary_output, + include_delta_base64=bool(args.include_delta_base64), + initial_playout_buffer_ms=initial_playout_buffer_ms, + ) + + +class JsonlWriter: + def __init__(self, output: Path) -> None: + output.parent.mkdir(parents=True, exist_ok=True) + self._handle = output.open("w", encoding="utf-8") + self._lock = asyncio.Lock() + + async def write(self, record: dict[str, Any]) -> None: + async with self._lock: + self._handle.write( + json.dumps(record, ensure_ascii=False, separators=(",", ":")) + "\n" + ) + self._handle.flush() + + def close(self) -> None: + self._handle.close() + + +async def receive_message( + ws: aiohttp.ClientWebSocketResponse, timeout_s: float | None +) -> aiohttp.WSMessage: + if timeout_s is None: + return await ws.receive() + return await asyncio.wait_for(ws.receive(), timeout=timeout_s) + + +def event_base( + *, run_id: str, call_id: str, call_number: int, sequence: int, started_at: float +) -> dict[str, Any]: + return { + "run_id": run_id, + "call_id": call_id, + "call_number": call_number, + "sequence": sequence, + "received_at_utc": utc_now(), + "received_after_call_start_ms": elapsed_ms(started_at), + } + + +async def capture_call( + *, + session: aiohttp.ClientSession, + config: Config, + writer: JsonlWriter, + run_id: str, + call_number: int, + call_id: str, + limiter: asyncio.Semaphore, + prewarmed_ws: aiohttp.ClientWebSocketResponse | None = None, +) -> str: + async with limiter: + started_at = time.perf_counter() + sequence = 0 + status = "unknown" + ws = prewarmed_ws + + async def write(kind: str, **fields: Any) -> None: + nonlocal sequence + sequence += 1 + record = event_base( + run_id=run_id, + call_id=call_id, + call_number=call_number, + sequence=sequence, + started_at=started_at, + ) + record["event"] = kind + record.update(fields) + await writer.write(record) + + try: + if ws is None: + await write("client.connecting", websocket_url=config.websocket_url) + connect_started_at = time.perf_counter() + ws = await session.ws_connect( + build_url(config), + headers={"Authorization": f"Bearer {config.api_key}"}, + heartbeat=20, + ) + await write( + "client.connected", connect_ms=elapsed_ms(connect_started_at) + ) + else: + await write("client.using_prewarmed_connection") + + await ws.send_json({"type": "text.clear"}) + await write("client.sent", message_type="text.clear") + clear_acknowledged = False + while not clear_acknowledged: + message = await receive_message(ws, config.event_timeout_s) + now = time.perf_counter() + if message.type != aiohttp.WSMsgType.TEXT: + await write( + "transport.message", + ws_type=message.type.name, + data=str(message.data), + ) + raise RuntimeError( + f"Expected audio.clear, received WebSocket {message.type.name}." + ) + try: + payload = json.loads(str(message.data)) + except json.JSONDecodeError: + await write("provider.invalid_json", raw_text=str(message.data)) + raise + await write( + "provider.event", + message_type=str(payload.get("type") or ""), + provider_message=payload, + received_monotonic_ms=round(now * 1000), + ) + clear_acknowledged = payload.get("type") == "audio.clear" + if payload.get("type") == "error": + raise RuntimeError(str(payload.get("message") or payload)) + + await ws.send_json({"type": "text.delta", "delta": config.text}) + await write( + "client.sent", message_type="text.delta", text_chars=len(config.text) + ) + await ws.send_json({"type": "text.done"}) + await write("client.sent", message_type="text.done") + + while True: + try: + message = await receive_message(ws, config.event_timeout_s) + except asyncio.TimeoutError: + await write( + "client.receive_timeout", + timeout_ms=None + if config.event_timeout_s is None + else round(config.event_timeout_s * 1000), + ) + status = "receive_timeout" + break + + now = time.perf_counter() + if message.type != aiohttp.WSMsgType.TEXT: + await write( + "transport.message", + ws_type=message.type.name, + data=str(message.data), + ) + status = f"transport_{message.type.name.lower()}" + break + + try: + payload = json.loads(str(message.data)) + except json.JSONDecodeError: + await write("provider.invalid_json", raw_text=str(message.data)) + status = "invalid_json" + break + + message_type = str(payload.get("type") or "") + if message_type == "audio.delta": + encoded = str(payload.get("delta") or "") + try: + pcm = base64.b64decode(encoded, validate=True) + except (binascii.Error, ValueError) as exc: + await write( + "provider.audio_delta_invalid", + message_type=message_type, + base64_characters=len(encoded), + error=str(exc), + provider_message={ + key: value + for key, value in payload.items() + if key != "delta" + }, + ) + status = "invalid_audio_delta" + break + record: dict[str, Any] = { + "message_type": message_type, + "base64_characters": len(encoded), + "audio_bytes": len(pcm), + "audio_ms": round( + len(pcm) * 1000 / (SAMPLE_RATE * PCM_BYTES_PER_SAMPLE), 3 + ), + "empty_audio": not bool(pcm), + "audio_sha256": hashlib.sha256(pcm).hexdigest(), + "provider_message": { + key: value + for key, value in payload.items() + if key != "delta" + }, + "received_monotonic_ms": round(now * 1000), + } + if config.include_delta_base64: + record["delta_base64"] = encoded + await write("provider.audio_delta", **record) + continue + + await write( + "provider.event", + message_type=message_type, + provider_message=payload, + received_monotonic_ms=round(now * 1000), + ) + if message_type == "audio.done": + status = "audio_done" + break + if message_type == "error": + status = "provider_error" + break + except Exception as exc: + if status == "unknown": + status = "exception" + await write( + "client.exception", error_type=type(exc).__name__, error=str(exc) + ) + finally: + if ws is not None and not ws.closed: + await ws.close() + await write("client.closed") + await write( + "client.finished", status=status, elapsed_ms=elapsed_ms(started_at) + ) + return status + + +async def prewarm_connection( + *, + session: aiohttp.ClientSession, + config: Config, + writer: JsonlWriter, + run_id: str, + call_number: int, + call_id: str, +) -> aiohttp.ClientWebSocketResponse: + started_at = time.perf_counter() + await writer.write( + { + **event_base( + run_id=run_id, + call_id=call_id, + call_number=call_number, + sequence=-1, + started_at=started_at, + ), + "event": "client.prewarm_connecting", + "websocket_url": config.websocket_url, + } + ) + try: + ws = await session.ws_connect( + build_url(config), + headers={"Authorization": f"Bearer {config.api_key}"}, + heartbeat=20, + ) + except Exception as exc: + await writer.write( + { + **event_base( + run_id=run_id, + call_id=call_id, + call_number=call_number, + sequence=0, + started_at=started_at, + ), + "event": "client.prewarm_failed", + "error_type": type(exc).__name__, + "error": str(exc), + } + ) + raise + await writer.write( + { + **event_base( + run_id=run_id, + call_id=call_id, + call_number=call_number, + sequence=0, + started_at=started_at, + ), + "event": "client.prewarm_ready", + "connect_ms": elapsed_ms(started_at), + } + ) + return ws + + +def percentile(values: list[float], percentile_value: float) -> float | None: + """Return the nearest-rank percentile without introducing a dependency.""" + if not values: + return None + ordered = sorted(values) + index = round((len(ordered) - 1) * percentile_value) + return ordered[max(0, min(index, len(ordered) - 1))] + + +def round_ms(value: float | None) -> float | None: + return None if value is None else round(value, 3) + + +def replay_underflow( + audio_frames: list[tuple[float, float]], initial_buffer_ms: float +) -> tuple[int, float, float]: + """Replay raw arrivals; with a zero prebuffer, playout starts at first audio.delta.""" + buffered_ms = 0.0 + playout_started = False + previous_arrival_ms: float | None = None + events = 0 + total_underflow_ms = 0.0 + largest_underflow_ms = 0.0 + + for arrived_at_ms, audio_ms in audio_frames: + if playout_started: + assert previous_arrival_ms is not None + elapsed_ms = arrived_at_ms - previous_arrival_ms + missing_audio_ms = max(0.0, elapsed_ms - buffered_ms) + if missing_audio_ms: + events += 1 + total_underflow_ms += missing_audio_ms + largest_underflow_ms = max(largest_underflow_ms, missing_audio_ms) + buffered_ms = max(0.0, buffered_ms - elapsed_ms) + audio_ms + previous_arrival_ms = arrived_at_ms + continue + + buffered_ms += audio_ms + if buffered_ms >= initial_buffer_ms: + playout_started = True + previous_arrival_ms = arrived_at_ms + + return events, total_underflow_ms, largest_underflow_ms + + +def metric_distribution(values: list[float]) -> dict[str, float | int | None]: + return { + "count": len(values), + "p50": round_ms(percentile(values, 0.50)), + "p95": round_ms(percentile(values, 0.95)), + "max": round_ms(max(values)) if values else None, + } + + +def read_finite_float(record: dict[str, Any], field: str) -> float | None: + """Read a finite numeric field without allowing a malformed capture to break analysis.""" + try: + value = float(record[field]) + except (KeyError, TypeError, ValueError): + return None + return value if math.isfinite(value) else None + + +def event_sequence(record: dict[str, Any]) -> int: + """Sort malformed or missing sequence values before regular events.""" + try: + return int(record.get("sequence", -2)) + except (TypeError, ValueError): + return -2 + + +def display_metric(value: float | None, unit: str = "") -> str: + return "n/a" if value is None else f"{value}{unit}" + + +def summarize_capture(config: Config, run_id: str) -> dict[str, Any]: + """Build a provider-facing summary exclusively from the raw JSONL capture.""" + calls: dict[str, list[tuple[int, dict[str, Any]]]] = {} + run_events: list[dict[str, Any]] = [] + analysis_warnings: list[str] = [] + with config.output.open(encoding="utf-8") as handle: + for line_number, line in enumerate(handle, start=1): + if not line.strip(): + continue + try: + record = json.loads(line) + except json.JSONDecodeError as exc: + analysis_warnings.append( + f"Skipped invalid JSONL record at line {line_number}: {exc.msg}." + ) + continue + if not isinstance(record, dict): + analysis_warnings.append( + f"Skipped non-object JSONL record at line {line_number}." + ) + continue + call_id = record.get("call_id") + if call_id: + calls.setdefault(str(call_id), []).append((line_number, record)) + else: + run_events.append(record) + + status_counts: Counter[str] = Counter() + error_counts: Counter[str] = Counter() + ttfb_ms: list[float] = [] + rtf: list[float] = [] + audio_duration_ms: list[float] = [] + frame_duration_ms: list[float] = [] + frame_gaps_ms: list[float] = [] + per_call_max_gaps_ms: list[float] = [] + largest_gap_observation: dict[str, Any] | None = None + frames_per_call: list[float] = [] + calls_with_audio = 0 + calls_with_underflow = 0 + underflow_events = 0 + total_underflow_ms = 0.0 + largest_underflow_ms = 0.0 + empty_audio_deltas = 0 + + for call_id, call_records in calls.items(): + call_records.sort(key=lambda item: event_sequence(item[1])) + records = [record for _, record in call_records] + for record in records: + event = str(record.get("event") or "") + message_type = str(record.get("message_type") or "") + if event == "client.finished": + status_counts[str(record.get("status") or "unknown")] += 1 + if event == "client.prewarm_failed": + error_counts["prewarm_failed"] += 1 + elif event == "client.exception": + error_counts["client_exception"] += 1 + elif event == "client.receive_timeout": + error_counts["receive_timeout"] += 1 + elif event == "provider.audio_delta_invalid": + error_counts["invalid_audio_delta"] += 1 + elif event == "provider.invalid_json": + error_counts["invalid_json"] += 1 + elif event == "transport.message": + error_counts[ + f"transport_{str(record.get('ws_type') or 'unknown').lower()}" + ] += 1 + elif event == "provider.event" and message_type == "error": + error_counts["provider_error"] += 1 + elif event == "provider.audio_delta" and record.get("empty_audio"): + empty_audio_deltas += 1 + + text_done_at_ms = None + for line_number, record in call_records: + if ( + record.get("event") != "client.sent" + or record.get("message_type") != "text.done" + ): + continue + text_done_at_ms = read_finite_float(record, "received_after_call_start_ms") + if text_done_at_ms is None: + analysis_warnings.append( + f"Skipped invalid text.done timestamp at line {line_number} for {call_id}." + ) + break + audio_frame_records: list[dict[str, Any]] = [] + for line_number, record in call_records: + if record.get("event") != "provider.audio_delta": + continue + audio_bytes = read_finite_float(record, "audio_bytes") + received_at_ms = read_finite_float(record, "received_after_call_start_ms") + frame_audio_ms = read_finite_float(record, "audio_ms") + if audio_bytes is None or received_at_ms is None or frame_audio_ms is None: + analysis_warnings.append( + f"Skipped malformed audio.delta at line {line_number} for {call_id}." + ) + continue + if audio_bytes > 0 and frame_audio_ms >= 0: + audio_frame_records.append(record) + audio_frames = [ + (float(record["received_after_call_start_ms"]), float(record["audio_ms"])) + for record in audio_frame_records + ] + if text_done_at_ms is None or not audio_frames: + continue + + calls_with_audio += 1 + frames_per_call.append(float(len(audio_frames))) + duration_ms = sum(frame[1] for frame in audio_frames) + audio_duration_ms.append(duration_ms) + frame_duration_ms.extend(frame[1] for frame in audio_frames) + ttfb_ms.append(audio_frames[0][0] - text_done_at_ms) + if duration_ms: + rtf.append((audio_frames[-1][0] - text_done_at_ms) / duration_ms) + gaps = [ + audio_frames[index][0] - audio_frames[index - 1][0] + for index in range(1, len(audio_frames)) + ] + frame_gaps_ms.extend(gaps) + if gaps: + per_call_max_gaps_ms.append(max(gaps)) + largest_gap_index = max( + range(1, len(audio_frames)), + key=lambda index: audio_frames[index][0] - audio_frames[index - 1][0], + ) + largest_gap_ms = ( + audio_frames[largest_gap_index][0] + - audio_frames[largest_gap_index - 1][0] + ) + if ( + largest_gap_observation is None + or largest_gap_ms > largest_gap_observation["gap_ms"] + ): + previous_record = audio_frame_records[largest_gap_index - 1] + current_record = audio_frame_records[largest_gap_index] + largest_gap_observation = { + "gap_ms": round_ms(largest_gap_ms), + "call_id": call_id, + "previous_audio_delta_sequence": previous_record.get("sequence"), + "previous_arrived_after_call_start_ms": audio_frames[ + largest_gap_index - 1 + ][0], + "previous_frame_audio_ms": audio_frames[largest_gap_index - 1][1], + "next_audio_delta_sequence": current_record.get("sequence"), + "next_arrived_after_call_start_ms": audio_frames[largest_gap_index][ + 0 + ], + "next_frame_audio_ms": audio_frames[largest_gap_index][1], + } + events, total_ms, largest_ms = replay_underflow( + audio_frames, config.initial_playout_buffer_ms + ) + if events: + calls_with_underflow += 1 + underflow_events += events + total_underflow_ms += total_ms + largest_underflow_ms = max(largest_underflow_ms, largest_ms) + + run_failed = next( + (record for record in run_events if record.get("event") == "run.failed"), None + ) + completed_calls = status_counts.get("audio_done", 0) + observed_errors = sum(error_counts.values()) + return { + "summary_version": 1, + "language": "en", + "source_events_jsonl": str(config.output), + "run_id": run_id, + "configuration": { + "requested_calls": config.calls, + "concurrency": config.concurrency, + "prewarm": config.prewarm, + "initial_playout_buffer_ms": config.initial_playout_buffer_ms, + "event_timeout_ms": None + if config.event_timeout_s is None + else round(config.event_timeout_s * 1000), + "connect_timeout_ms": None + if config.connect_timeout_s is None + else round(config.connect_timeout_s * 1000), + "pcm_format": f"{SAMPLE_RATE} Hz, mono, signed 16-bit PCM", + }, + "completion": { + "calls_observed": len(calls), + "calls_with_audio": calls_with_audio, + "audio_done": completed_calls, + "audio_done_rate": round(completed_calls / config.calls, 4) + if config.calls + else None, + "finished_statuses": dict(sorted(status_counts.items())), + "run_failed": bool(run_failed), + "run_failure": None if run_failed is None else run_failed.get("error"), + }, + "latency": { + "ttfb_ms": metric_distribution(ttfb_ms), + "rtf": metric_distribution(rtf), + "generated_audio_ms_per_call": metric_distribution(audio_duration_ms), + }, + "frame_delivery": { + "raw_interframe_arrival_gap_definition": "Timestamp difference between consecutive non-empty provider.audio_delta messages for the same call. It excludes TTFB and does not include either frame's audio duration.", + "nonempty_audio_frames_per_call": metric_distribution(frames_per_call), + "audio_ms_per_frame": metric_distribution(frame_duration_ms), + "raw_interframe_arrival_gap_ms": metric_distribution(frame_gaps_ms), + "per_call_max_raw_interframe_arrival_gap_ms": metric_distribution( + per_call_max_gaps_ms + ), + "largest_raw_interframe_arrival_gap": largest_gap_observation, + "empty_audio_delta_count": empty_audio_deltas, + }, + "simulated_playout": { + "definition": "TTFB is excluded. The replay clock starts when the first non-empty audio.delta arrives. If initial_buffer_ms is greater than zero, playout is intentionally delayed until that much PCM has accumulated; otherwise it starts immediately at the first frame.", + "initial_buffer_ms": config.initial_playout_buffer_ms, + "calls_with_underflow": calls_with_underflow, + "calls_with_underflow_rate": round(calls_with_underflow / config.calls, 4) + if config.calls + else None, + "underflow_events": underflow_events, + "total_underflow_ms": round_ms(total_underflow_ms), + "largest_underflow_ms": round_ms(largest_underflow_ms), + }, + "errors": { + "total_observed": observed_errors, + "by_type": dict(sorted(error_counts.items())), + "empty_audio_delta_count": empty_audio_deltas, + }, + "analysis_warnings": analysis_warnings, + } + + +def format_summary(summary: dict[str, Any]) -> str: + """Create a concise, copy-pasteable English summary for terminal output.""" + configuration = summary["configuration"] + completion = summary["completion"] + latency = summary["latency"] + frames = summary["frame_delivery"] + playout = summary["simulated_playout"] + errors = summary["errors"] + warnings = summary["analysis_warnings"] + ttfb = latency["ttfb_ms"] + rtf = latency["rtf"] + gaps = frames["raw_interframe_arrival_gap_ms"] + largest_gap = frames["largest_raw_interframe_arrival_gap"] + largest_gap_context = ( + "n/a" + if largest_gap is None + else ( + f"call {largest_gap['call_id']}, sequences " + f"{largest_gap['previous_audio_delta_sequence']}->{largest_gap['next_audio_delta_sequence']}" + ) + ) + return "\n".join( + [ + "TTS WebSocket Capture Summary", + f"Load: {configuration['requested_calls']} calls at C={configuration['concurrency']} (prewarm: {'on' if configuration['prewarm'] else 'off'})", + f"Completion: {completion['audio_done']}/{configuration['requested_calls']} audio.done ({completion['audio_done_rate']:.1%})", + "TTFB (text.done -> first non-empty audio): " + f"p50 {display_metric(ttfb['p50'], ' ms')} | p95 {display_metric(ttfb['p95'], ' ms')} | max {display_metric(ttfb['max'], ' ms')}", + "RTF (last audio arrival / generated audio): " + f"p50 {display_metric(rtf['p50'])} | p95 {display_metric(rtf['p95'])} | max {display_metric(rtf['max'])}", + "Raw inter-frame arrival gap (consecutive non-empty audio.delta; excludes TTFB and frame duration): " + f"p50 {display_metric(gaps['p50'], ' ms')} | p95 {display_metric(gaps['p95'], ' ms')} | max {display_metric(gaps['max'], ' ms')} ({largest_gap_context})", + "Audio payload duration per frame: " + f"p50 {display_metric(frames['audio_ms_per_frame']['p50'], ' ms')} | p95 {display_metric(frames['audio_ms_per_frame']['p95'], ' ms')} | max {display_metric(frames['audio_ms_per_frame']['max'], ' ms')}", + "Simulated underflow " + f"(TTFB excluded; {playout['initial_buffer_ms']} ms intentional prebuffer): {playout['calls_with_underflow']}/{configuration['requested_calls']} calls | {playout['underflow_events']} events | largest {display_metric(playout['largest_underflow_ms'], ' ms')} | total {display_metric(playout['total_underflow_ms'], ' ms')}", + f"Errors: {errors['total_observed']} observed | empty audio.delta frames: {errors['empty_audio_delta_count']} | by type: {errors['by_type'] or 'none'}", + f"Analysis warnings: {len(warnings)}" + + (f" | {warnings[0]}" if warnings else ""), + ] + ) + + +async def run(argv: list[str] | None = None) -> int: + config = load_config(parser().parse_args(argv)) + run_id = uuid.uuid4().hex + writer = JsonlWriter(config.output) + limiter = asyncio.Semaphore(config.concurrency) + statuses: list[str] = [] + call_ids = { + number: f"call-{number:04d}-{uuid.uuid4().hex[:8]}" + for number in range(1, config.calls + 1) + } + try: + if config.prewarm and config.calls != config.concurrency: + raise ValueError( + "--prewarm requires --calls to equal --concurrency so the measured load window is consistent." + ) + await writer.write( + { + "event": "run.started", + "run_id": run_id, + "started_at_utc": utc_now(), + "calls": config.calls, + "concurrency": config.concurrency, + "prewarm": config.prewarm, + "event_timeout_ms": None + if config.event_timeout_s is None + else round(config.event_timeout_s * 1000), + "connect_timeout_ms": None + if config.connect_timeout_s is None + else round(config.connect_timeout_s * 1000), + "include_delta_base64": config.include_delta_base64, + "pcm_format": { + "sample_rate": SAMPLE_RATE, + "channels": 1, + "sample_width_bytes": PCM_BYTES_PER_SAMPLE, + }, + } + ) + timeout = aiohttp.ClientTimeout( + total=None, sock_connect=config.connect_timeout_s + ) + connector = aiohttp.TCPConnector(limit=config.concurrency) + async with aiohttp.ClientSession( + timeout=timeout, connector=connector + ) as session: + prewarmed: dict[int, aiohttp.ClientWebSocketResponse | None] = { + number: None for number in call_ids + } + if config.prewarm: + prewarm_started_at = time.perf_counter() + results = await asyncio.gather( + *[ + prewarm_connection( + session=session, + config=config, + writer=writer, + run_id=run_id, + call_number=number, + call_id=call_ids[number], + ) + for number in call_ids + ], + return_exceptions=True, + ) + failures = [item for item in results if isinstance(item, BaseException)] + if failures: + for item in results: + if ( + isinstance(item, aiohttp.ClientWebSocketResponse) + and not item.closed + ): + await item.close() + raise RuntimeError( + f"Prewarm failed for {len(failures)}/{config.calls} connections." + ) + prewarmed = {number: results[number - 1] for number in call_ids} + await writer.write( + { + "event": "run.prewarm_finished", + "run_id": run_id, + "finished_at_utc": utc_now(), + "prewarm_duration_ms": elapsed_ms(prewarm_started_at), + "connections_ready": len(prewarmed), + } + ) + statuses = await asyncio.gather( + *[ + capture_call( + session=session, + config=config, + writer=writer, + run_id=run_id, + call_number=call_number, + call_id=call_ids[call_number], + limiter=limiter, + prewarmed_ws=prewarmed[call_number], + ) + for call_number in call_ids + ] + ) + await writer.write( + { + "event": "run.finished", + "run_id": run_id, + "finished_at_utc": utc_now(), + "statuses": { + status: statuses.count(status) for status in sorted(set(statuses)) + }, + } + ) + except BaseException as exc: + await writer.write( + { + "event": "run.failed", + "run_id": run_id, + "failed_at_utc": utc_now(), + "error_type": type(exc).__name__, + "error": str(exc), + } + ) + raise + finally: + writer.close() + try: + summary = summarize_capture(config, run_id) + config.summary_output.parent.mkdir(parents=True, exist_ok=True) + config.summary_output.write_text( + json.dumps(summary, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + print(format_summary(summary)) + print(f"Summary JSON: {config.summary_output}") + except Exception as exc: + print( + f"Unable to generate summary JSON: {type(exc).__name__}: {exc}", + file=sys.stderr, + ) + return 0 if statuses and all(status == "audio_done" for status in statuses) else 1 + + +def main() -> None: + raise SystemExit(asyncio.run(run())) + + +if __name__ == "__main__": + main() diff --git a/src/app/utils/audio_backlog.py b/src/app/utils/audio_backlog.py new file mode 100644 index 0000000..ade198f --- /dev/null +++ b/src/app/utils/audio_backlog.py @@ -0,0 +1,111 @@ +"""Shed de backlog de audio guiado por energia, compartilhado entre a entrada +(cliente -> LiveKit) e a saida (agente -> cliente) do bridge. + +Ideia central: quando o backlog de uma fila de frames PCM16 passa de um limiar, +em vez de descartar cegamente os frames mais antigos -- o que pode cortar fala -- +preferimos descartar apenas os frames silenciosos, comprimindo pausas sem remover +amostras de voz. O descarte cego fica reservado para o regime de rajada da +entrada, onde o excesso e grande e a prioridade e recuperar latencia. Na saida +(fala do agente ja gerada, esperando ser tocada) o descarte cego nunca deve ser +usado: so silencio pode sair. +""" + +from __future__ import annotations + +import audioop +from dataclasses import dataclass +from typing import Any + +_FULL_SCALE = 32768.0 + + +def rms_threshold_from_dbfs(dbfs: float) -> int: + """Converte um limiar em dBFS para amplitude RMS PCM16 (ex.: -50 -> ~103).""" + try: + return max(0, int(_FULL_SCALE * (10 ** (float(dbfs) / 20.0)))) + except (ValueError, OverflowError): + return 0 + + +def frames_from_ms(value_ms: int, frame_ms: int, *, minimum: int = 1) -> int: + value_ms = max(0, int(value_ms)) + frame_ms = max(1, int(frame_ms)) + return max(minimum, (value_ms + frame_ms - 1) // frame_ms) + + +def frame_rms(frame: bytes, sample_width: int = 2) -> int: + try: + return audioop.rms(frame, sample_width) + except Exception: + return 0 + + +def is_silent_frame(frame: bytes, rms_threshold: int, *, sample_width: int = 2) -> bool: + return frame_rms(frame, sample_width) <= rms_threshold + + +@dataclass(frozen=True) +class ShedResult: + dropped: int = 0 + dropped_silent: int = 0 + dropped_voiced: int = 0 + mode: str = "none" + + +def shed_queue_backlog( + queue: Any, + *, + keep_frames: int, + rms_threshold: int, + blind: bool, + sample_width: int = 2, +) -> ShedResult: + """Reduz ``queue`` (``asyncio.Queue[bytes]``) ate ``keep_frames``. + + - ``blind=True``: descarta os frames mais antigos ate chegar em + ``keep_frames`` (regime de rajada -- recupera latencia mesmo cortando + conteudo). + - ``blind=False`` (guiado por energia): drena a fila e descarta apenas frames + silenciosos (rms <= ``rms_threshold``), preservando a fala. Pode terminar + acima de ``keep_frames`` se nao houver silencio suficiente -- de proposito. + + IMPORTANTE: deve ser chamada sem ``await`` entre os get/put. No event loop + isso e atomico: nenhum produtor/consumidor da fila roda no meio da drenagem. + """ + keep_frames = max(0, int(keep_frames)) + qsize = queue.qsize() + need = qsize - keep_frames + if need <= 0: + return ShedResult() + + if blind: + dropped = 0 + while queue.qsize() > keep_frames: + try: + queue.get_nowait() + except Exception: + break + dropped += 1 + return ShedResult(dropped=dropped, dropped_voiced=dropped, mode="blind") + + frames: list[bytes] = [] + while True: + try: + frames.append(queue.get_nowait()) + except Exception: + break + + to_drop = need + dropped_silent = 0 + kept: list[bytes] = [] + for frame in frames: + if to_drop > 0 and is_silent_frame(frame, rms_threshold, sample_width=sample_width): + to_drop -= 1 + dropped_silent += 1 + continue + kept.append(frame) + + for frame in kept: + queue.put_nowait(frame) + + return ShedResult(dropped=dropped_silent, dropped_silent=dropped_silent, mode="energy") diff --git a/src/app/utils/background.py b/src/app/utils/background.py new file mode 100644 index 0000000..c0aba2b --- /dev/null +++ b/src/app/utils/background.py @@ -0,0 +1,524 @@ +from __future__ import annotations + +import asyncio +import math +import logging +import os +import time +import audioop +from dataclasses import dataclass +from typing import Any, Callable, Mapping, Optional +from fastapi import WebSocket, WebSocketDisconnect + +from app.utils.audio_backlog import ( + frames_from_ms, + rms_threshold_from_dbfs, + shed_queue_backlog, +) +from app.utils.logging import log_flow_event + +logger = logging.getLogger("background") + + +def _out_env_int(name: str, default: int) -> int: + try: + return int(os.getenv(name, str(default))) + except (TypeError, ValueError): + return default + + +def _out_env_float(name: str, default: float) -> float: + try: + return float(os.getenv(name, str(default))) + except (TypeError, ValueError): + return default + + +def _out_env_bool(name: str, default: bool) -> bool: + raw = os.getenv(name) + if raw is None: + return default + return raw.strip().lower() in {"1", "true", "yes", "on"} + + +@dataclass +class AudioActivity: + last_agent_out: float # "última fala com energia" do agente (idle/mock-stop) + + +@dataclass +class BridgeOutputStats: + total_ws_frames: int = 0 + agent_audio_frames: int = 0 + silence_frames: int = 0 + agent_audio_bursts: int = 0 + first_agent_audio_sent: bool = False + last_agent_audio_sent_at: float = 0.0 + + +@dataclass(frozen=True) +class AudioOutputBacklogConfig: + """Instrumentacao e shed do caminho de saida (agente -> cliente). + + O backlog da agent_q e, em boa parte, fala do agente ja gerada esperando ser + tocada em tempo real -- por isso o shed de saida e SEMPRE guiado por energia + (so silencio sai; nunca corta voz). O descarte cego da entrada nao se aplica. + """ + + metrics_enabled: bool = True + latency_alert_ms: int = 1000 + latency_log_interval_s: float = 15.0 + pacing_debt_threshold_ms: int = 200 + shed_enabled: bool = True + shed_threshold_ms: int = 400 + shed_keep_ms: int = 200 + shed_check_interval_frames: int = 5 + silence_dbfs: float = -60.0 + + def shed_threshold_frames(self, frame_ms: int) -> int: + return frames_from_ms(self.shed_threshold_ms, frame_ms) + + def shed_keep_frames(self, frame_ms: int) -> int: + return min( + self.shed_threshold_frames(frame_ms), + frames_from_ms(self.shed_keep_ms, frame_ms), + ) + + @property + def silence_rms_threshold(self) -> int: + return rms_threshold_from_dbfs(self.silence_dbfs) + + +def audio_output_backlog_config_from_env() -> AudioOutputBacklogConfig: + return AudioOutputBacklogConfig( + metrics_enabled=_out_env_bool("AUDIO_OUT_LATENCY_METRICS_ENABLED", True), + latency_alert_ms=max(0, _out_env_int("AUDIO_OUT_LATENCY_ALERT_MS", 1000)), + latency_log_interval_s=max( + 0.1, _out_env_float("AUDIO_OUT_LATENCY_LOG_INTERVAL_S", 15.0) + ), + pacing_debt_threshold_ms=max( + 0, _out_env_int("AUDIO_OUT_PACING_DEBT_THRESHOLD_MS", 200) + ), + shed_enabled=_out_env_bool("AUDIO_OUT_BACKLOG_SHED_ENABLED", True), + shed_threshold_ms=max(0, _out_env_int("AUDIO_OUT_BACKLOG_SHED_THRESHOLD_MS", 400)), + shed_keep_ms=max(0, _out_env_int("AUDIO_OUT_BACKLOG_SHED_KEEP_MS", 200)), + shed_check_interval_frames=max( + 1, _out_env_int("AUDIO_OUT_BACKLOG_SHED_CHECK_INTERVAL_FRAMES", 5) + ), + silence_dbfs=_out_env_float("AUDIO_OUT_BACKLOG_SILENCE_DBFS", -60.0), + ) + + +class AudioOutputLatencyTracker: + """Relogio + metrica + shed do caminho de saida, espelhando a entrada. + + - queue_ms/peak_queue_ms: quanto audio do agente espera na agent_q. + - pacing_debt_dropped_ms: divida de pacing jogada fora quando o loop atrasa + (antes era um resync mudo ``tick = now``, sem contador). + - total_dropped_ms: silencio descartado pelo shed de saida. + """ + + def __init__( + self, + config: AudioOutputBacklogConfig, + *, + frame_ms: int, + flow_logger: Any, + timeline: Any = None, + call_context: Optional[Mapping[str, Any]] = None, + debug_event_publisher: Optional[Callable[[str, Mapping[str, Any]], None]] = None, + ) -> None: + self.config = config + self.frame_ms = max(1, int(frame_ms)) + self.flow_logger = flow_logger + self.timeline = timeline + self.call_context = dict(call_context or {}) + self.debug_event_publisher = debug_event_publisher + self.peak_queue_ms = 0 + self.total_dropped_frames = 0 + self.total_dropped_silent = 0 + self.pacing_debt_events = 0 + self.pacing_debt_dropped_ms = 0 + self.last_log_at: Optional[float] = None + self.alert_active = False + + def _emit(self, step: str, payload: Mapping[str, Any]) -> None: + log_flow_event(self.flow_logger, step, **payload) + if self.timeline is not None: + try: + self.timeline.emit(step, **payload) + except Exception: + pass + + def note_pacing_debt(self, debt_ms: int, *, now: Optional[float] = None) -> None: + debt_ms = max(0, int(debt_ms)) + if debt_ms <= 0: + return + self.pacing_debt_events += 1 + self.pacing_debt_dropped_ms += debt_ms + if self.config.metrics_enabled: + self._emit( + "audio_out_pacing_debt", + { + **self.call_context, + "debt_ms": debt_ms, + "pacing_debt_events": self.pacing_debt_events, + "pacing_debt_dropped_ms": self.pacing_debt_dropped_ms, + }, + ) + + def record_shed( + self, + *, + dropped_frames: int, + dropped_silent: int, + mode: str, + queue_before_frames: int, + queue_after_frames: int, + ) -> None: + if dropped_frames <= 0: + return + self.total_dropped_frames += int(dropped_frames) + self.total_dropped_silent += int(dropped_silent) + if not self.config.metrics_enabled: + return + payload = { + **self.call_context, + "mode": mode, + "content_policy": "silence_only" if mode == "energy" else "oldest_frames", + "sync_action": "compress_silence" if mode == "energy" else "drop_audio_to_catch_up", + "dropped_frames": int(dropped_frames), + "dropped_ms": int(dropped_frames) * self.frame_ms, + "dropped_silent_frames": int(dropped_silent), + "queue_before_ms": int(queue_before_frames) * self.frame_ms, + "queue_after_ms": int(queue_after_frames) * self.frame_ms, + "total_dropped_ms": self.total_dropped_frames * self.frame_ms, + "shed_threshold_ms": self.config.shed_threshold_ms, + "shed_keep_ms": self.config.shed_keep_ms, + } + self._emit("audio_out_latency_shed", payload) + if self.debug_event_publisher is not None: + self.debug_event_publisher("bridge.audio_out.shed", payload) + + def maybe_log( + self, + *, + queue_frames: int, + now: float, + reason: str, + force: bool = False, + ) -> None: + queue_ms = max(0, int(queue_frames)) * self.frame_ms + self.peak_queue_ms = max(self.peak_queue_ms, queue_ms) + if not self.config.metrics_enabled: + return + alert_active = ( + self.config.latency_alert_ms > 0 + and queue_ms >= self.config.latency_alert_ms + ) + alert_started = alert_active and not self.alert_active + self.alert_active = alert_active + if self.last_log_at is None: + self.last_log_at = now + if not force and not alert_started: + return + should_log = ( + force + or alert_started + or (now - self.last_log_at) >= self.config.latency_log_interval_s + ) + if not should_log: + return + self.last_log_at = now + self._emit( + "audio_out_latency", + { + **self.call_context, + "queue_frames": max(0, int(queue_frames)), + "queue_ms": queue_ms, + "peak_queue_ms": self.peak_queue_ms, + "pacing_debt_events": self.pacing_debt_events, + "pacing_debt_dropped_ms": self.pacing_debt_dropped_ms, + "total_dropped_ms": self.total_dropped_frames * self.frame_ms, + "reason": reason, + "alert_ms": self.config.latency_alert_ms, + }, + ) +def _config_value(value: Any) -> str: + if value in (None, ""): + return "" + return str(value).strip() + + +def _gain_from_linear(value: Any) -> Optional[float]: + text = _config_value(value) + if not text: + return None + try: + gain = float(text) + return max(0.0, min(10.0, gain)) if math.isfinite(gain) else None + except Exception: + return None + + +def _get_ws_output_gain( + *, + output_gain: Any = None, +) -> float: + """ + Prioriza callConfig.ws.outputGain. Fallback: WS_OUTPUT_GAIN; default 1.0. + """ + for gain in ( + _gain_from_linear(output_gain), + _gain_from_linear(os.getenv("WS_OUTPUT_GAIN", "")), + ): + if gain is not None: + return gain + return 1.0 + + +def _apply_pcm_gain(frame: bytes, sample_width: int, gain: float) -> bytes: + if gain == 1.0: + return frame + + try: + return audioop.mul(frame, sample_width, gain) + except Exception: + return frame + + + +def _agent_speech_rms_threshold() -> int: + """Energia mínima para eventos de início de fala; não filtra o áudio.""" + try: + dbfs = float(os.getenv("AGENT_OUT_SPEECH_DBFS", "-52")) + except (TypeError, ValueError): + dbfs = -52.0 + if not math.isfinite(dbfs): + dbfs = -52.0 + return int(32768 * (10 ** (max(-90.0, min(0.0, dbfs)) / 20.0))) + + + +async def ws_out_loop( + ws: WebSocket, + agent_q: asyncio.Queue[bytes], + *, + frame_ms: int, + bytes_per_frame: int, + first_agent_audio_sent: Optional[asyncio.Event] = None, + flow_logger: Any = None, + timeline: Any = None, + call_context: Optional[Mapping[str, Any]] = None, + output_stats: Optional[BridgeOutputStats] = None, + output_gain: Any = None, + recording_sink: Any = None, + debug_event_publisher: Optional[Callable[[str, Mapping[str, Any]], None]] = None, +): + frame_interval = frame_ms / 1000.0 + silence = b"\x00" * bytes_per_frame + flow_logger = flow_logger or logger + flow_context = dict(call_context or {}) + output_stats = output_stats or BridgeOutputStats() + + audio_burst_gap_s = float(os.getenv("FLOW_AUDIO_BURST_GAP_S", "0.8")) + ws_output_gain = _get_ws_output_gain( + output_gain=output_gain, + ) + agent_speech_rms_threshold = _agent_speech_rms_threshold() + if ws_output_gain != 1.0: + logger.info("[ws-out] aplicando ganho global nos frames enviados ao WS gain=%.3f", ws_output_gain) + + # relogio + metrica + shed da saida (espelha a instrumentacao da entrada) + out_config = audio_output_backlog_config_from_env() + out_tracker = AudioOutputLatencyTracker( + out_config, + frame_ms=frame_ms, + flow_logger=flow_logger, + timeline=timeline, + call_context=flow_context, + debug_event_publisher=debug_event_publisher, + ) + out_shed_threshold_frames = out_config.shed_threshold_frames(frame_ms) + out_shed_keep_frames = out_config.shed_keep_frames(frame_ms) + out_silence_rms = out_config.silence_rms_threshold + pacing_debt_threshold_s = out_config.pacing_debt_threshold_ms / 1000.0 + if out_config.metrics_enabled: + out_tracker._emit( + "audio_out_backlog_config", + { + **flow_context, + "shed_enabled": out_config.shed_enabled, + "shed_threshold_ms": out_config.shed_threshold_ms, + "shed_keep_ms": out_config.shed_keep_ms, + "silence_dbfs": out_config.silence_dbfs, + "latency_alert_ms": out_config.latency_alert_ms, + "pacing_debt_threshold_ms": out_config.pacing_debt_threshold_ms, + "shed_policy": "silence_only", + }, + ) + + def _emit_flow(step: str, **fields: Any) -> None: + log_flow_event(flow_logger, step, **flow_context, **fields) + + tick = time.monotonic() + n = 0 + + while True: + now = time.monotonic() + + if now < tick: + await asyncio.sleep(tick - now) + now = time.monotonic() + + if now - tick > pacing_debt_threshold_s: + # loop atrasou: em vez de resync mudo, contabiliza a divida de pacing. + out_tracker.note_pacing_debt(round((now - tick) * 1000), now=now) + tick = now + + # 0) shed de backlog da saida: so silencio sai, nunca corta a fala do + # agente. Rate-limited para nao redrenar a fila a cada tick. + if ( + out_config.shed_enabled + and (n % out_config.shed_check_interval_frames) == 0 + ): + out_queue_before = agent_q.qsize() + if out_queue_before > out_shed_threshold_frames: + out_shed = shed_queue_backlog( + agent_q, + keep_frames=out_shed_keep_frames, + rms_threshold=out_silence_rms, + blind=False, + ) + if out_shed.dropped: + out_tracker.record_shed( + dropped_frames=out_shed.dropped, + dropped_silent=out_shed.dropped_silent, + mode=out_shed.mode, + queue_before_frames=out_queue_before, + queue_after_frames=agent_q.qsize(), + ) + + # 1) Consome exatamente um frame por tick. O frame e encaminhado como + # veio; silencio temporal do agente so e removido pelo shed acima. + agent_frame: Optional[bytes] = None + agent_frame_rms = 0 + agent_frame_has_speech = False + frame_source = "silence" + try: + cand = agent_q.get_nowait() + except asyncio.QueueEmpty: + cand = None + + if cand is not None: + if len(cand) != bytes_per_frame: + if len(cand) > bytes_per_frame: + cand = cand[:bytes_per_frame] + else: + cand = cand + (b"\x00" * (bytes_per_frame - len(cand))) + agent_frame = cand + try: + agent_frame_rms = audioop.rms(cand, 2) + except Exception: + agent_frame_rms = 0 + agent_frame_has_speech = agent_frame_rms > agent_speech_rms_threshold + if agent_frame is not None: + frame = agent_frame + frame_source = "agent" + else: + # fila vazia: preenche o pacing de 20ms com silêncio + frame = silence + + frame = _apply_pcm_gain(frame, 2, ws_output_gain) + + try: + await ws.send_bytes(frame) + except WebSocketDisconnect: + return + except RuntimeError as exc: + if "close message has been sent" in str(exc): + return + raise + + if recording_sink is not None: + try: + recording_sink.record_output_frame(frame) + except Exception: + pass + + output_stats.total_ws_frames += 1 + if frame_source == "agent": + output_stats.agent_audio_frames += 1 + new_burst = ( + output_stats.last_agent_audio_sent_at <= 0.0 + or (now - output_stats.last_agent_audio_sent_at) >= audio_burst_gap_s + ) + if agent_frame_has_speech: + output_stats.last_agent_audio_sent_at = now + + if agent_frame_has_speech and not output_stats.first_agent_audio_sent: + output_stats.first_agent_audio_sent = True + _emit_flow( + "audio_to_tia_first_frame", + destination="ws_client", + ws_frame=output_stats.total_ws_frames, + agent_audio_frames=output_stats.agent_audio_frames, + queued_agent_frames=agent_q.qsize(), + rms=agent_frame_rms, + bytes=len(frame), + ) + if timeline is not None: + timeline.emit( + "agent_audio_first_frame_sent_to_tia", + **flow_context, + destination="ws_client", + ws_frame=output_stats.total_ws_frames, + agent_audio_frames=output_stats.agent_audio_frames, + queued_agent_frames=agent_q.qsize(), + rms=agent_frame_rms, + bytes=len(frame), + ) + + if agent_frame_has_speech and new_burst: + output_stats.agent_audio_bursts += 1 + _emit_flow( + "audio_to_tia", + destination="ws_client", + seq=output_stats.agent_audio_bursts, + ws_frame=output_stats.total_ws_frames, + agent_audio_frames=output_stats.agent_audio_frames, + queued_agent_frames=agent_q.qsize(), + rms=agent_frame_rms, + bytes=len(frame), + ) + if timeline is not None: + timeline.emit( + "agent_audio_sent_to_tia", + **flow_context, + destination="ws_client", + seq=output_stats.agent_audio_bursts, + ws_frame=output_stats.total_ws_frames, + agent_audio_frames=output_stats.agent_audio_frames, + queued_agent_frames=agent_q.qsize(), + rms=agent_frame_rms, + bytes=len(frame), + ) + else: + output_stats.silence_frames += 1 + + if ( + agent_frame is not None + and agent_frame_has_speech + and first_agent_audio_sent is not None + and not first_agent_audio_sent.is_set() + ): + first_agent_audio_sent.set() + + out_tracker.maybe_log( + queue_frames=agent_q.qsize(), + now=now, + reason="tick", + ) + + n += 1 + tick += frame_interval diff --git a/src/app/utils/call_timeline.py b/src/app/utils/call_timeline.py new file mode 100644 index 0000000..ce40244 --- /dev/null +++ b/src/app/utils/call_timeline.py @@ -0,0 +1,356 @@ +from __future__ import annotations + +import atexit +import json +import logging +import os +import queue +import threading +import time +from collections import defaultdict +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Mapping + +try: + import fcntl +except ImportError: # pragma: no cover - indisponível no Windows + fcntl = None + + +def _sanitize_for_filename(value: str) -> str: + raw = (value or "").strip() + if not raw: + return "unknown" + + allowed = [] + for ch in raw: + if ch.isalnum() or ch in {"-", "_", "."}: + allowed.append(ch) + else: + allowed.append("_") + + sanitized = "".join(allowed).strip("._") + return sanitized or "unknown" + + +def _json_safe(value: Any) -> Any: + if value is None: + return None + if isinstance(value, (str, int, float, bool)): + return value + if isinstance(value, Path): + return str(value) + if isinstance(value, Mapping): + return {str(key): _json_safe(item) for key, item in value.items()} + if isinstance(value, (list, tuple, set)): + return [_json_safe(item) for item in value] + return str(value) + + +def _bounded_env_int(name: str, default: int, *, minimum: int, maximum: int) -> int: + try: + value = int(os.getenv(name, str(default)) or str(default)) + except (TypeError, ValueError): + value = default + return max(minimum, min(maximum, value)) + + +class _WriterCommand: + __slots__ = ("kind", "event") + + def __init__(self, kind: str) -> None: + self.kind = kind + self.event = threading.Event() + + +class _TimelineWriter: + """Serializa timelines numa thread dedicada, fora do event loop.""" + + def __init__( + self, + max_queue: int, + *, + warning_interval_s: float = 60.0, + ) -> None: + self._queue: queue.Queue[tuple[Path, str] | _WriterCommand] = queue.Queue( + maxsize=max(1, max_queue) + ) + self._thread: threading.Thread | None = None + self._lifecycle_lock = threading.Lock() + self._stats_lock = threading.Lock() + self._closed = False + self._dropped = 0 + self._write_errors = 0 + self._last_drop_warning_at = 0.0 + self._last_error_warning_at = 0.0 + self._warning_interval_s = max(1.0, float(warning_interval_s)) + self._logger = logging.getLogger(f"{__name__}.writer") + + @property + def dropped(self) -> int: + with self._stats_lock: + return self._dropped + + @property + def write_errors(self) -> int: + with self._stats_lock: + return self._write_errors + + def start(self) -> bool: + with self._lifecycle_lock: + if self._closed: + return False + if self._thread is not None and self._thread.is_alive(): + return True + self._thread = threading.Thread( + target=self._run, + name="call-timeline-writer", + daemon=True, + ) + self._thread.start() + return True + + def submit(self, path: Path, line: str) -> bool: + if not self.start(): + self._record_drop(reason="writer_closed") + return False + reason = "" + with self._lifecycle_lock: + if self._closed: + reason = "writer_closed" + else: + try: + self._queue.put_nowait((path, line)) + return True + except queue.Full: + reason = "queue_full" + self._record_drop(reason=reason) + return False + + def flush(self, timeout: float | None = 2.0) -> bool: + with self._lifecycle_lock: + if self._closed: + return self._queue.empty() + if not self.start(): + return False + command = _WriterCommand("flush") + deadline = None if timeout is None else time.monotonic() + max(0.0, timeout) + with self._lifecycle_lock: + if self._closed: + return self._queue.empty() + if not self._enqueue_command(command, deadline): + return False + remaining = None if deadline is None else max(0.0, deadline - time.monotonic()) + return command.event.wait(remaining) + + def shutdown(self, timeout: float | None = 2.0) -> bool: + with self._lifecycle_lock: + if self._closed: + thread = self._thread + return thread is None or not thread.is_alive() + self._closed = True + thread = self._thread + + if thread is None: + return True + + command = _WriterCommand("shutdown") + deadline = None if timeout is None else time.monotonic() + max(0.0, timeout) + if not self._enqueue_command(command, deadline): + return False + remaining = None if deadline is None else max(0.0, deadline - time.monotonic()) + if not command.event.wait(remaining): + return False + remaining = None if deadline is None else max(0.0, deadline - time.monotonic()) + thread.join(remaining) + return not thread.is_alive() + + def _enqueue_command( + self, + command: _WriterCommand, + deadline: float | None, + ) -> bool: + try: + if deadline is None: + self._queue.put(command) + else: + self._queue.put(command, timeout=max(0.0, deadline - time.monotonic())) + return True + except queue.Full: + return False + + def _run(self) -> None: + while True: + first = self._queue.get() + batch: list[tuple[Path, str]] = [] + commands: list[_WriterCommand] = [] + self._collect(first, batch, commands) + while True: + try: + item = self._queue.get_nowait() + except queue.Empty: + break + self._collect(item, batch, commands) + + if batch: + self._write_batch(batch) + + should_stop = False + for command in commands: + command.event.set() + should_stop = should_stop or command.kind == "shutdown" + if should_stop: + return + + @staticmethod + def _collect( + item: tuple[Path, str] | _WriterCommand, + batch: list[tuple[Path, str]], + commands: list[_WriterCommand], + ) -> None: + if isinstance(item, _WriterCommand): + commands.append(item) + else: + batch.append(item) + + def _write_batch(self, batch: list[tuple[Path, str]]) -> None: + by_path: dict[Path, list[str]] = defaultdict(list) + for path, line in batch: + by_path[path].append(line) + for path, lines in by_path.items(): + try: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("a", encoding="utf-8") as handle: + if fcntl is not None: + fcntl.flock(handle.fileno(), fcntl.LOCK_EX) + try: + handle.write("".join(f"{line}\n" for line in lines)) + finally: + if fcntl is not None: + fcntl.flock(handle.fileno(), fcntl.LOCK_UN) + except Exception as exc: + self._record_write_error(path=path, exc=exc) + + def _record_drop(self, *, reason: str) -> None: + now = time.monotonic() + with self._stats_lock: + self._dropped += 1 + dropped = self._dropped + should_warn = now - self._last_drop_warning_at >= self._warning_interval_s + if should_warn: + self._last_drop_warning_at = now + if should_warn: + self._logger.warning( + "CALL_TIMELINE_ASYNC_DROP | reason=%s | dropped_total=%s | queue_size=%s", + reason, + dropped, + self._queue.qsize(), + ) + + def _record_write_error(self, *, path: Path, exc: Exception) -> None: + now = time.monotonic() + with self._stats_lock: + self._write_errors += 1 + errors = self._write_errors + should_warn = now - self._last_error_warning_at >= self._warning_interval_s + if should_warn: + self._last_error_warning_at = now + if should_warn: + self._logger.warning( + "CALL_TIMELINE_ASYNC_WRITE_FAIL | path=%s | error=%s: %s | errors_total=%s", + path, + type(exc).__name__, + exc, + errors, + ) + + +_WRITER = _TimelineWriter( + _bounded_env_int( + "CALL_TIMELINE_QUEUE_MAX", + 10_000, + minimum=1, + maximum=1_000_000, + ), + warning_interval_s=_bounded_env_int( + "ASYNC_IO_WARNING_INTERVAL_S", + 60, + minimum=1, + maximum=3_600, + ), +) +atexit.register(_WRITER.shutdown) + + +class CallTimeline: + def __init__( + self, + *, + logger: logging.Logger, + component: str, + timeline_id: str, + protocol: str = "", + room: str = "", + session_id: str = "", + phone_number: str = "", + origin_unix_ms: int | None = None, + ) -> None: + self._logger = logger + self._component = (component or "unknown").strip().lower() + self._timeline_id = _sanitize_for_filename(timeline_id or session_id or protocol or room) + self._protocol = (protocol or "").strip() + self._room = (room or "").strip() + self._session_id = (session_id or "").strip() + self._phone_number = (phone_number or "").strip() + self._origin_unix_ms = int(origin_unix_ms or round(time.time() * 1000)) + self._enabled = os.getenv("CALL_TIMELINE_ENABLED", "1") == "1" + self._console_enabled = os.getenv("CALL_TIMELINE_CONSOLE", "0") == "1" + self._dir = Path(os.getenv("CALL_TIMELINE_DIR", "./timeline")) + self._path = self._dir / f"{self._timeline_id}.jsonl" + if self._enabled: + _WRITER.start() + + @property + def origin_unix_ms(self) -> int: + return self._origin_unix_ms + + @property + def path(self) -> Path: + return self._path + + def emit(self, event: str, **fields: Any) -> None: + if not self._enabled: + return + + now_ms = round(time.time() * 1000) + record = { + "ts": datetime.now(timezone.utc).isoformat(), + "t_rel_ms": max(0, int(now_ms - self._origin_unix_ms)), + "component": self._component, + "event": (event or "unknown").strip(), + "timeline_id": self._timeline_id, + "protocol": self._protocol, + "room": self._room, + "session_id": self._session_id, + "phone_number": self._phone_number, + } + for key, value in fields.items(): + record[str(key)] = _json_safe(value) + + line = json.dumps(record, ensure_ascii=False) + self._append_line(line) + + if self._console_enabled: + self._logger.info("TIMELINE | %s", line) + + def _append_line(self, line: str) -> None: + _WRITER.submit(self._path, line) + + def flush(self, timeout: float | None = 2.0) -> bool: + """Aguarda eventos já enfileirados; não deve ser chamado no loop de áudio.""" + return _WRITER.flush(timeout=timeout) + + @property + def dropped_events(self) -> int: + return _WRITER.dropped diff --git a/src/app/utils/dump.py b/src/app/utils/dump.py new file mode 100644 index 0000000..a6c1289 --- /dev/null +++ b/src/app/utils/dump.py @@ -0,0 +1,76 @@ +import os +import time +import wave +import uuid +import logging +from pathlib import Path + +logger = logging.getLogger(__name__) + +def _dump_pcm16_wav( + pcm16: bytes, + *, + sample_rate: int = 16000, + channels: int = 1, + dump_dir: str | None = None, + prefix: str = "stt", +) -> str | None: + """ + Salva PCM16LE (raw) em WAV para inspeção. + Retorna path salvo ou None se desabilitado. + """ + dump_dir = dump_dir or os.getenv("STT_DUMP_DIR", "") + if not dump_dir: + return None + + Path(dump_dir).mkdir(parents=True, exist_ok=True) + ts = time.strftime("%Y%m%dT%H%M%S") + fname = f"{prefix}_{ts}_{uuid.uuid4().hex[:8]}.wav" + path = str(Path(dump_dir) / fname) + + with wave.open(path, "wb") as w: + w.setnchannels(channels) + w.setsampwidth(2) # PCM16 => 2 bytes + w.setframerate(sample_rate) + w.writeframes(pcm16) + + return path + + +def _dump_bytes(path: str, b: bytes) -> None: + Path(path).parent.mkdir(parents=True, exist_ok=True) + Path(path).write_bytes(b) + +def _dump_audio_for_debug( + *, + req_id: str, + pcm: bytes, + wav_bytes: bytes, + sample_rate: int, + channels: int, +) -> dict: + """ + Salva PCM e WAV em disco quando STT_DUMP_DIR estiver setado. + Retorna dict com paths (ou vazio se desabilitado). + """ + dump_dir = os.getenv("STT_DUMP_DIR", "").strip() + if not dump_dir: + return {} + + ts = time.strftime("%Y%m%dT%H%M%S") + base = str(Path(dump_dir) / f"stt_{ts}_{req_id}") + + pcm_path = base + ".pcm" # raw PCM16LE + wav_path = base + ".wav" # WAV pronto (exatamente o que vai pro STT) + + _dump_bytes(pcm_path, pcm) + _dump_bytes(wav_path, wav_bytes) + + return { + "pcm_path": pcm_path, + "wav_path": wav_path, + "bytes_pcm": len(pcm), + "bytes_wav": len(wav_bytes), + "sample_rate": sample_rate, + "channels": channels, + } \ No newline at end of file diff --git a/src/app/utils/export.py b/src/app/utils/export.py new file mode 100644 index 0000000..f20e584 --- /dev/null +++ b/src/app/utils/export.py @@ -0,0 +1,191 @@ +import json +import csv +import os +import re +from datetime import datetime +from typing import Union, List, Dict, Any, Optional + + +def _only_digits(value: Any) -> str: + return re.sub(r"\D+", "", str(value or "")) + + +def _find_first(obj: Any, keys: List[str]) -> Optional[Any]: + """ + Procura recursivamente (dict/list) pelo primeiro campo existente em `keys`. + Retorna o valor encontrado ou None. + """ + if isinstance(obj, dict): + for k in keys: + if k in obj and obj[k] not in (None, ""): + return obj[k] + for v in obj.values(): + found = _find_first(v, keys) + if found not in (None, ""): + return found + elif isinstance(obj, list): + for it in obj: + found = _find_first(it, keys) + if found not in (None, ""): + return found + return None + + +def _parse_date_time(date_str: Any, time_str: Any) -> Optional[datetime]: + date_s = str(date_str).strip() + time_s = str(time_str).strip() + + # Normaliza hora + if re.fullmatch(r"\d{6}", time_s): # HHMMSS + time_s = f"{time_s[0:2]}:{time_s[2:4]}:{time_s[4:6]}" + elif re.fullmatch(r"\d{4}", time_s): # HHMM + time_s = f"{time_s[0:2]}:{time_s[2:4]}:00" + elif re.fullmatch(r"\d{2}:\d{2}$", time_s): # HH:MM + time_s = f"{time_s}:00" + + # Normaliza data (suporta 2026-01-26, 26/01/2026, 20260126) + parsed_date = None + for fmt in ("%Y-%m-%d", "%d/%m/%Y", "%Y%m%d"): + try: + parsed_date = datetime.strptime(date_s, fmt).date() + break + except ValueError: + pass + + if not parsed_date: + return None + + try: + parsed_time = datetime.strptime(time_s, "%H:%M:%S").time() + except ValueError: + return None + + return datetime.combine(parsed_date, parsed_time) + + +def _parse_iso_datetime(value: Any) -> Optional[datetime]: + if value in (None, ""): + return None + s = str(value).strip() + try: + # suporta "Z" no final + dt = datetime.fromisoformat(s.replace("Z", "+00:00")) + # remove timezone para ficar consistente no filename + return dt.replace(tzinfo=None) + except ValueError: + return None + + +def _extract_success_purchase(payload: Any) -> Dict[str, Any]: + """ + Encontra o dict success_purchase (quando existir) no payload. + """ + if isinstance(payload, dict): + sp = payload.get("success_purchase") + return sp if isinstance(sp, dict) else {} + if isinstance(payload, list) and payload and isinstance(payload[0], dict): + sp = payload[0].get("success_purchase") + return sp if isinstance(sp, dict) else {} + return {} + + +def _build_filename(payload: Any, uuid: Any) -> str: + # Telefone: tenta vários campos comuns (success_purchase / formatted_call / etc.) + phone_val = _find_first( + payload, + ["NUM_TELEFONE", "NUM_ACESSO", "telefone", "phone", "phone_number"], + ) + phone = _only_digits(phone_val) or _only_digits(uuid) or "unknown" + + # Data/hora: tenta DATA + HORA primeiro + date_val = _find_first(payload, ["DATA", "data", "date"]) + time_val = _find_first(payload, ["HORA", "hora", "time"]) + + dt = None + if date_val and time_val: + dt = _parse_date_time(date_val, time_val) + + # Fallback: tenta dataISO / timestamp ISO + if not dt: + dt = _parse_iso_datetime( + _find_first(payload, ["dataISO", "DATA_ISO", "timestamp", "datetime"]) + ) + + # Último fallback: agora + if not dt: + dt = datetime.now() + + stamp = dt.strftime("%Y%m%dT%H%M%S") + return f"{phone}_{stamp}.csv" + + +def json_to_csv( + json_input: Union[str, List[Dict[str, Any]], Dict[str, Any]], + uuid, + encoding: str = "utf-8", +) -> None: + # 1) Carrega o JSON (arquivo ou objeto em memória) + if isinstance(json_input, str): + with open(json_input, "r", encoding=encoding) as f: + payload = json.load(f) + else: + payload = json_input + + print("Json de retorno", payload) + + # 2) Valida success_purchase (mantendo sua regra) + success_purchase = _extract_success_purchase(payload) + if not success_purchase: + return + + # 3) Nome do arquivo: telefone_data + fname = _build_filename(payload, uuid) + + # Base da pasta vinda da env, padrão ./recordings + base_dir = os.getenv("EXPORT_DIR", "./recordings") + + # 4) Salva direto em ./recordings (sem subpasta do telefone) + os.makedirs(base_dir, exist_ok=True) + path = os.path.join(base_dir, fname) + + # 5) Normaliza os dados para lista de dicts (para achatar) + data = payload + if isinstance(data, dict): + data = [data] + elif not isinstance(data, list): + raise ValueError("O JSON deve ser uma lista de objetos ou um único objeto.") + + if not data: + raise ValueError("JSON está vazio, não há dados para escrever no CSV.") + + # 6) Achata cada linha + linhas_flat: List[Dict[str, Any]] = [] + fieldnames_set = set() + + for item in data: + if not isinstance(item, dict): + raise ValueError("Cada item da lista JSON deve ser um objeto (dict).") + + flat_row: Dict[str, Any] = {} + + for key, value in item.items(): + if isinstance(value, dict): + for subkey, subvalue in value.items(): + flat_row[subkey] = subvalue + else: + flat_row[key] = value + + linhas_flat.append(flat_row) + fieldnames_set.update(flat_row.keys()) + + # Opcional: ordena as colunas para ficar estável + fieldnames = sorted(fieldnames_set) + + # 7) Escreve CSV + with open(path, "w", encoding=encoding, newline="") as csvfile: + writer = csv.DictWriter(csvfile, fieldnames=fieldnames) + writer.writeheader() + for row in linhas_flat: + writer.writerow(row) + + print(f"CSV gerado com sucesso em: {path}") diff --git a/src/app/utils/full_call_recording.py b/src/app/utils/full_call_recording.py new file mode 100644 index 0000000..e8da323 --- /dev/null +++ b/src/app/utils/full_call_recording.py @@ -0,0 +1,903 @@ +from __future__ import annotations + +import asyncio +import hashlib +import logging +import os +import queue +import re +import threading +import time +import uuid +import wave +import weakref +from collections import deque +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Any, Mapping + +from app.utils.logging import log_flow_event +from app.utils.stt_audio_upload import ( + OCIUploadConfig, + _OCI_UPLOAD_LOCK, + _oci_client, + _oci_upload_config_from_env, + _reset_oci_client, +) + +logger = logging.getLogger(__name__) + +_DEFAULT_QUEUE_SIZE = 1024 +_DEFAULT_MAX_CLIENT_BUFFER_FRAMES = 512 +_UPLOAD_QUEUE_SIZE = 8 +_UPLOAD_MAX_ATTEMPTS = 3 +_UPLOAD_RETRY_BASE_DELAY_S = 0.25 +_MAX_PART_CHARS = 64 +_SAFE_PART_RE = re.compile(r"[^A-Za-z0-9_.-]+") +_SENTINEL = object() +_RESERVED_LOG_FIELD_NAMES = frozenset( + { + "audio_duration_ms", + "bytes", + "duration_ms", + "dropped_client_buffer_frames", + "dropped_queue_frames", + "error", + "frame_ms", + "frames", + "object_name", + "path", + "reason", + "sample_rate", + "session_id", + } +) + + +@dataclass(frozen=True, slots=True) +class EntireCallRecordingResult: + object_name: str + path: Path + bytes: int + frames: int + duration_ms: int + dropped_queue_frames: int + dropped_client_buffer_frames: int + upload_enqueued: bool + + +@dataclass(frozen=True, slots=True) +class _RecordingEvent: + source: str + frame: bytes + + +@dataclass(frozen=True, slots=True) +class EntireCallUploadItem: + object_name: str + path: Path + session_id: str + bytes: int + duration_ms: int + dropped_queue_frames: int + dropped_client_buffer_frames: int + metadata: Mapping[str, Any] + logger: Any | None = None + timeline: Any | None = None + + +def _safe_path_part(value: Any, *, fallback: str) -> str: + raw = str(value or "").strip() + if not raw: + raw = fallback + safe = _SAFE_PART_RE.sub("_", raw).strip("._-") or fallback + if len(safe) <= _MAX_PART_CHARS: + return safe + + digest = hashlib.sha1(safe.encode("utf-8")).hexdigest()[:12] + prefix_len = max(1, _MAX_PART_CHARS - len(digest) - 1) + return f"{safe[:prefix_len].rstrip('._-')}-{digest}" + + +def build_entire_call_object_name( + *, + session_id: Any, + now: datetime | None = None, +) -> str: + dt = now or datetime.now() + session_part = _safe_path_part(session_id, fallback="unknown-session") + return f"{dt:%Y-%m-%d}/entire_call/{session_part}.wav" + + +def _env_bool(name: str, default: bool) -> bool: + raw = os.getenv(name) + if raw is None: + return default + return raw.strip().lower() in {"1", "true", "yes", "on"} + + +def _tmp_dir_from_env() -> Path: + return Path(os.getenv("ENTIRE_CALL_RECORDING_TMP_DIR", "./recordings/entire_call_tmp")) + + +def _metadata_fields(metadata: Mapping[str, Any]) -> dict[str, Any]: + return { + str(key): value + for key, value in metadata.items() + if value not in (None, "") and str(key) not in _RESERVED_LOG_FIELD_NAMES + } + + +def _emit_timeline(timeline: Any | None, event: str, **fields: Any) -> None: + if timeline is None: + return + try: + timeline.emit(event, **fields) + except Exception: + pass + + +def _cleanup_path(path: Path, active_logger: Any | None = None) -> None: + try: + path.unlink(missing_ok=True) + except Exception: + (active_logger or logger).debug( + "[entire-call-recording] failed cleaning temp file path=%s", + path, + exc_info=True, + ) + + +def _put_object_file_sync(item: EntireCallUploadItem, config: OCIUploadConfig) -> Any: + with _OCI_UPLOAD_LOCK: + client = _oci_client(config) + with item.path.open("rb") as body: + return client.put_object( + namespace_name=config.namespace, + bucket_name=config.bucket, + object_name=item.object_name, + put_object_body=body, + ) + + +def _upload_max_attempts() -> int: + raw = os.getenv("ENTIRE_CALL_RECORDING_UPLOAD_MAX_ATTEMPTS", str(_UPLOAD_MAX_ATTEMPTS)) + try: + return max(1, int(raw)) + except (TypeError, ValueError): + return _UPLOAD_MAX_ATTEMPTS + + +def _upload_retry_base_delay_s() -> float: + raw = os.getenv( + "ENTIRE_CALL_RECORDING_UPLOAD_RETRY_BASE_DELAY_S", + str(_UPLOAD_RETRY_BASE_DELAY_S), + ) + try: + return max(0.0, float(raw)) + except (TypeError, ValueError): + return _UPLOAD_RETRY_BASE_DELAY_S + + +async def _upload_item(item: EntireCallUploadItem) -> None: + started = time.perf_counter() + config = _oci_upload_config_from_env() + if config is None: + raise RuntimeError("OCI Object Storage upload config is missing") + + max_attempts = _upload_max_attempts() + base_delay_s = _upload_retry_base_delay_s() + response: Any | None = None + for attempt in range(1, max_attempts + 1): + try: + response = await asyncio.to_thread(_put_object_file_sync, item, config) + break + except Exception as exc: + # Do not close/reset the shared HTTP session while an STT upload is + # using it. The same lock serializes reset and put_object calls. + with _OCI_UPLOAD_LOCK: + _reset_oci_client() + if attempt >= max_attempts: + raise + delay_s = base_delay_s * (2 ** (attempt - 1)) + log_flow_event( + item.logger or logger, + "entire_call_recording_upload_retry", + session_id=item.session_id, + object_name=item.object_name, + attempt=attempt, + max_attempts=max_attempts, + delay_ms=round(delay_s * 1000), + error_type=type(exc).__name__, + **_metadata_fields(item.metadata), + ) + if delay_s: + await asyncio.sleep(delay_s) + + if response is None: + raise RuntimeError("OCI Object Storage upload returned no response") + duration_ms = round((time.perf_counter() - started) * 1000) + active_logger = item.logger or logger + log_flow_event( + active_logger, + "entire_call_recording_upload_done", + session_id=item.session_id, + object_name=item.object_name, + bytes=item.bytes, + audio_duration_ms=item.duration_ms, + duration_ms=duration_ms, + bucket=config.bucket, + namespace=config.namespace, + oci_status=getattr(response, "status", None), + dropped_queue_frames=item.dropped_queue_frames, + dropped_client_buffer_frames=item.dropped_client_buffer_frames, + **_metadata_fields(item.metadata), + ) + _emit_timeline( + item.timeline, + "entire_call_recording_upload_done", + session_id=item.session_id, + object_name=item.object_name, + bytes=item.bytes, + audio_duration_ms=item.duration_ms, + duration_ms=duration_ms, + bucket=config.bucket, + namespace=config.namespace, + oci_status=getattr(response, "status", None), + dropped_queue_frames=item.dropped_queue_frames, + dropped_client_buffer_frames=item.dropped_client_buffer_frames, + **_metadata_fields(item.metadata), + ) + + +class _EntireCallUploadWorker: + def __init__(self, loop: asyncio.AbstractEventLoop) -> None: + self._loop = loop + self._queue: asyncio.Queue[EntireCallUploadItem] = asyncio.Queue(maxsize=_UPLOAD_QUEUE_SIZE) + self._task: asyncio.Task[None] | None = None + + def enqueue(self, item: EntireCallUploadItem) -> bool: + if self._task is None or self._task.done(): + self._task = self._loop.create_task( + self._run(), + name="entire-call-recording-upload-worker", + ) + try: + self._queue.put_nowait(item) + return True + except asyncio.QueueFull: + return False + + async def _run(self) -> None: + while True: + item = await self._queue.get() + active_logger = item.logger or logger + try: + await _upload_item(item) + _cleanup_path(item.path, active_logger) + except Exception: + log_flow_event( + active_logger, + "entire_call_recording_upload_failed", + session_id=item.session_id, + object_name=item.object_name, + bytes=item.bytes, + audio_duration_ms=item.duration_ms, + retained_path=str(item.path), + dropped_queue_frames=item.dropped_queue_frames, + dropped_client_buffer_frames=item.dropped_client_buffer_frames, + **_metadata_fields(item.metadata), + ) + _emit_timeline( + item.timeline, + "entire_call_recording_upload_failed", + session_id=item.session_id, + object_name=item.object_name, + bytes=item.bytes, + audio_duration_ms=item.duration_ms, + retained_path=str(item.path), + dropped_queue_frames=item.dropped_queue_frames, + dropped_client_buffer_frames=item.dropped_client_buffer_frames, + **_metadata_fields(item.metadata), + ) + active_logger.exception( + "[entire-call-recording][upload_failed] session_id=%s object_name=%s", + item.session_id, + item.object_name, + ) + finally: + self._queue.task_done() + + +_UPLOAD_WORKERS: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, _EntireCallUploadWorker] = ( + weakref.WeakKeyDictionary() +) +_UPLOAD_WORKERS_LOCK = threading.Lock() + + +def _worker_for_loop(loop: asyncio.AbstractEventLoop) -> _EntireCallUploadWorker: + with _UPLOAD_WORKERS_LOCK: + worker = _UPLOAD_WORKERS.get(loop) + if worker is None: + worker = _EntireCallUploadWorker(loop) + _UPLOAD_WORKERS[loop] = worker + return worker + + +def enqueue_entire_call_upload( + *, + path: Path, + object_name: str, + session_id: str, + bytes: int, + duration_ms: int, + dropped_queue_frames: int = 0, + dropped_client_buffer_frames: int = 0, + metadata: Mapping[str, Any] | None = None, + logger_override: Any | None = None, + timeline: Any | None = None, + cleanup_on_drop: bool = False, +) -> str | None: + active_logger = logger_override or logger + metadata = metadata or {} + if _oci_upload_config_from_env() is None: + log_flow_event( + active_logger, + "entire_call_recording_upload_drop", + session_id=session_id, + object_name=object_name, + reason="missing_oci_config", + bytes=bytes, + audio_duration_ms=duration_ms, + **_metadata_fields(metadata), + ) + _emit_timeline( + timeline, + "entire_call_recording_upload_drop", + session_id=session_id, + object_name=object_name, + reason="missing_oci_config", + bytes=bytes, + audio_duration_ms=duration_ms, + **_metadata_fields(metadata), + ) + if cleanup_on_drop: + _cleanup_path(path, active_logger) + return None + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + log_flow_event( + active_logger, + "entire_call_recording_upload_drop", + session_id=session_id, + object_name=object_name, + reason="no_running_loop", + bytes=bytes, + audio_duration_ms=duration_ms, + **_metadata_fields(metadata), + ) + if cleanup_on_drop: + _cleanup_path(path, active_logger) + return None + + item = EntireCallUploadItem( + object_name=object_name, + path=path, + session_id=session_id, + bytes=bytes, + duration_ms=duration_ms, + dropped_queue_frames=dropped_queue_frames, + dropped_client_buffer_frames=dropped_client_buffer_frames, + metadata=metadata, + logger=active_logger, + timeline=timeline, + ) + queued = _worker_for_loop(loop).enqueue(item) + if not queued: + log_flow_event( + active_logger, + "entire_call_recording_upload_drop", + session_id=session_id, + object_name=object_name, + reason="queue_full", + bytes=bytes, + audio_duration_ms=duration_ms, + **_metadata_fields(metadata), + ) + _emit_timeline( + timeline, + "entire_call_recording_upload_drop", + session_id=session_id, + object_name=object_name, + reason="queue_full", + bytes=bytes, + audio_duration_ms=duration_ms, + **_metadata_fields(metadata), + ) + if cleanup_on_drop: + _cleanup_path(path, active_logger) + return object_name + + log_flow_event( + active_logger, + "entire_call_recording_upload_enqueued", + session_id=session_id, + object_name=object_name, + bytes=bytes, + audio_duration_ms=duration_ms, + dropped_queue_frames=dropped_queue_frames, + dropped_client_buffer_frames=dropped_client_buffer_frames, + **_metadata_fields(metadata), + ) + _emit_timeline( + timeline, + "entire_call_recording_upload_enqueued", + session_id=session_id, + object_name=object_name, + bytes=bytes, + audio_duration_ms=duration_ms, + dropped_queue_frames=dropped_queue_frames, + dropped_client_buffer_frames=dropped_client_buffer_frames, + **_metadata_fields(metadata), + ) + return object_name + + +class EntireCallRecorder: + def __init__( + self, + *, + session_id: str, + sample_rate: int, + channels: int, + sample_width: int, + frame_ms: int, + bytes_per_frame: int, + tmp_dir: Path, + logger_override: Any | None = None, + timeline: Any | None = None, + metadata: Mapping[str, Any] | None = None, + queue_size: int = _DEFAULT_QUEUE_SIZE, + max_client_buffer_frames: int = _DEFAULT_MAX_CLIENT_BUFFER_FRAMES, + object_name: str | None = None, + ) -> None: + self.session_id = str(session_id or "").strip() or "unknown-session" + self.sample_rate = int(sample_rate) + self.source_channels = int(channels) + if self.source_channels != 1: + raise ValueError("entire call recorder requires mono source frames") + self.channels = 2 + self.sample_width = int(sample_width) + self.frame_ms = int(frame_ms) + self.bytes_per_frame = int(bytes_per_frame) + self.object_name = object_name or build_entire_call_object_name(session_id=self.session_id) + self._tmp_dir = tmp_dir + self._logger = logger_override or logger + self._timeline = timeline + self._metadata = dict(metadata or {}) + self._queue: queue.Queue[_RecordingEvent | object] = queue.Queue( + maxsize=max(1, int(queue_size)) + ) + self._max_client_buffer_frames = max(1, int(max_client_buffer_frames)) + self._thread: threading.Thread | None = None + self._state_lock = threading.Lock() + self._started = False + self._closed = False + self._finalized = False + self._thread_error: BaseException | None = None + self._frames_written = 0 + self._dropped_queue_frames = 0 + self._dropped_client_buffer_frames = 0 + self._path = self._tmp_dir / ( + f"{_safe_path_part(self.session_id, fallback='unknown-session')}-{uuid.uuid4().hex}.wav" + ) + + @property + def path(self) -> Path: + return self._path + + @property + def dropped_queue_frames(self) -> int: + return self._dropped_queue_frames + + def start(self) -> None: + with self._state_lock: + if self._started: + return + self._tmp_dir.mkdir(parents=True, exist_ok=True) + self._started = True + self._thread = threading.Thread( + target=self._run, + name=f"entire-call-recorder-{self.session_id[:16]}", + daemon=True, + ) + self._thread.start() + + log_flow_event( + self._logger, + "entire_call_recording_started", + session_id=self.session_id, + object_name=self.object_name, + path=str(self._path), + sample_rate=self.sample_rate, + channels=self.channels, + channel_left="client", + channel_right="agent", + frame_ms=self.frame_ms, + **_metadata_fields(self._metadata), + ) + _emit_timeline( + self._timeline, + "entire_call_recording_started", + session_id=self.session_id, + object_name=self.object_name, + path=str(self._path), + sample_rate=self.sample_rate, + channels=self.channels, + channel_left="client", + channel_right="agent", + frame_ms=self.frame_ms, + **_metadata_fields(self._metadata), + ) + + def record_client_frame(self, frame: bytes) -> None: + self._enqueue("client", frame) + + def record_output_frame(self, frame: bytes) -> None: + self._enqueue("output", frame) + + def _enqueue(self, source: str, frame: bytes) -> None: + try: + with self._state_lock: + if not self._started or self._closed: + return + self._queue.put_nowait(_RecordingEvent(source=source, frame=self._normalize_frame(frame))) + except queue.Full: + self._record_queue_drop(source=source) + except Exception: + self._logger.debug( + "[entire-call-recording] failed enqueueing frame source=%s", + source, + exc_info=True, + ) + + def _record_queue_drop(self, *, source: str) -> None: + with self._state_lock: + self._dropped_queue_frames += 1 + dropped = self._dropped_queue_frames + + if dropped == 1 or dropped % 100 == 0: + log_flow_event( + self._logger, + "entire_call_recording_drop", + session_id=self.session_id, + object_name=self.object_name, + source=source, + reason="queue_full", + dropped_queue_frames=dropped, + queue_size=self._queue.qsize(), + **_metadata_fields(self._metadata), + ) + _emit_timeline( + self._timeline, + "entire_call_recording_drop", + session_id=self.session_id, + object_name=self.object_name, + source=source, + reason="queue_full", + dropped_queue_frames=dropped, + queue_size=self._queue.qsize(), + **_metadata_fields(self._metadata), + ) + + def _normalize_frame(self, frame: bytes) -> bytes: + data = bytes(frame or b"") + if len(data) == self.bytes_per_frame: + return data + if len(data) > self.bytes_per_frame: + return data[: self.bytes_per_frame] + return data + (b"\x00" * (self.bytes_per_frame - len(data))) + + def _interleave_stereo_frames(self, client_frame: bytes, output_frame: bytes) -> bytes: + """Build stereo PCM with the client on the left and the agent on the right.""" + stereo = bytearray(len(client_frame) + len(output_frame)) + write_offset = 0 + for sample_offset in range(0, self.bytes_per_frame, self.sample_width): + next_offset = sample_offset + self.sample_width + stereo[write_offset : write_offset + self.sample_width] = client_frame[ + sample_offset:next_offset + ] + write_offset += self.sample_width + stereo[write_offset : write_offset + self.sample_width] = output_frame[ + sample_offset:next_offset + ] + write_offset += self.sample_width + return bytes(stereo) + + def _write_stereo_frame( + self, + wav_handle: wave.Wave_write, + *, + client_frame: bytes, + output_frame: bytes, + ) -> None: + wav_handle.writeframesraw( + self._interleave_stereo_frames(client_frame, output_frame) + ) + self._frames_written += 1 + + def _run(self) -> None: + silence = b"\x00" * self.bytes_per_frame + client_frames: deque[bytes] = deque() + try: + with wave.open(str(self._path), "wb") as wav_handle: + wav_handle.setnchannels(self.channels) + wav_handle.setsampwidth(self.sample_width) + wav_handle.setframerate(self.sample_rate) + + while True: + event = self._queue.get() + try: + if event is _SENTINEL: + break + + if not isinstance(event, _RecordingEvent): + continue + + if event.source == "client": + if len(client_frames) >= self._max_client_buffer_frames: + client_frames.popleft() + self._dropped_client_buffer_frames += 1 + client_frames.append(event.frame) + continue + + if event.source == "output": + client_frame = client_frames.popleft() if client_frames else silence + self._write_stereo_frame( + wav_handle, + client_frame=client_frame, + output_frame=event.frame, + ) + finally: + self._queue.task_done() + + while client_frames: + self._write_stereo_frame( + wav_handle, + client_frame=client_frames.popleft(), + output_frame=silence, + ) + except BaseException as exc: + self._thread_error = exc + self._logger.exception( + "[entire-call-recording][writer_failed] session_id=%s object_name=%s", + self.session_id, + self.object_name, + ) + + def _close_and_join(self) -> None: + with self._state_lock: + if not self._started or self._closed: + return + self._closed = True + + self._queue.put(_SENTINEL, timeout=5.0) + thread = self._thread + if thread is not None: + thread.join(timeout=5.0) + if thread.is_alive(): + raise TimeoutError("entire call recorder writer did not stop in time") + + async def finalize(self, *, enqueue_upload: bool = True) -> EntireCallRecordingResult | None: + with self._state_lock: + if self._finalized: + return None + self._finalized = True + + try: + await asyncio.to_thread(self._close_and_join) + except Exception: + log_flow_event( + self._logger, + "entire_call_recording_failed", + session_id=self.session_id, + object_name=self.object_name, + reason="writer_finalize_failed", + **_metadata_fields(self._metadata), + ) + _emit_timeline( + self._timeline, + "entire_call_recording_failed", + session_id=self.session_id, + object_name=self.object_name, + reason="writer_finalize_failed", + **_metadata_fields(self._metadata), + ) + self._logger.exception( + "[entire-call-recording][finalize_failed] session_id=%s object_name=%s", + self.session_id, + self.object_name, + ) + _cleanup_path(self._path, self._logger) + return None + + if self._thread_error is not None: + log_flow_event( + self._logger, + "entire_call_recording_failed", + session_id=self.session_id, + object_name=self.object_name, + reason="writer_failed", + error=type(self._thread_error).__name__, + **_metadata_fields(self._metadata), + ) + _emit_timeline( + self._timeline, + "entire_call_recording_failed", + session_id=self.session_id, + object_name=self.object_name, + reason="writer_failed", + error=type(self._thread_error).__name__, + **_metadata_fields(self._metadata), + ) + _cleanup_path(self._path, self._logger) + return None + + file_bytes = self._path.stat().st_size if self._path.exists() else 0 + duration_ms = round(self._frames_written * self.frame_ms) + result = EntireCallRecordingResult( + object_name=self.object_name, + path=self._path, + bytes=file_bytes, + frames=self._frames_written, + duration_ms=duration_ms, + dropped_queue_frames=self._dropped_queue_frames, + dropped_client_buffer_frames=self._dropped_client_buffer_frames, + upload_enqueued=False, + ) + log_flow_event( + self._logger, + "entire_call_recording_finalized", + session_id=self.session_id, + object_name=self.object_name, + path=str(self._path), + bytes=file_bytes, + frames=self._frames_written, + audio_duration_ms=duration_ms, + channels=self.channels, + channel_left="client", + channel_right="agent", + dropped_queue_frames=self._dropped_queue_frames, + dropped_client_buffer_frames=self._dropped_client_buffer_frames, + **_metadata_fields(self._metadata), + ) + _emit_timeline( + self._timeline, + "entire_call_recording_finalized", + session_id=self.session_id, + object_name=self.object_name, + path=str(self._path), + bytes=file_bytes, + frames=self._frames_written, + audio_duration_ms=duration_ms, + channels=self.channels, + channel_left="client", + channel_right="agent", + dropped_queue_frames=self._dropped_queue_frames, + dropped_client_buffer_frames=self._dropped_client_buffer_frames, + **_metadata_fields(self._metadata), + ) + + if not enqueue_upload: + return result + + queued_object = enqueue_entire_call_upload( + path=self._path, + object_name=self.object_name, + session_id=self.session_id, + bytes=file_bytes, + duration_ms=duration_ms, + dropped_queue_frames=self._dropped_queue_frames, + dropped_client_buffer_frames=self._dropped_client_buffer_frames, + metadata=self._metadata, + logger_override=self._logger, + timeline=self._timeline, + ) + return EntireCallRecordingResult( + object_name=self.object_name, + path=self._path, + bytes=file_bytes, + frames=self._frames_written, + duration_ms=duration_ms, + dropped_queue_frames=self._dropped_queue_frames, + dropped_client_buffer_frames=self._dropped_client_buffer_frames, + upload_enqueued=queued_object is not None, + ) + + +def create_entire_call_recorder_from_env( + *, + session_id: str, + sample_rate: int, + channels: int, + sample_width: int, + frame_ms: int, + bytes_per_frame: int, + logger_override: Any | None = None, + timeline: Any | None = None, + metadata: Mapping[str, Any] | None = None, +) -> EntireCallRecorder | None: + active_logger = logger_override or logger + metadata = metadata or {} + if not _env_bool("ENTIRE_CALL_RECORDING_ENABLED", True): + log_flow_event( + active_logger, + "entire_call_recording_disabled", + session_id=session_id, + reason="disabled_by_env", + **_metadata_fields(metadata), + ) + _emit_timeline( + timeline, + "entire_call_recording_disabled", + session_id=session_id, + reason="disabled_by_env", + **_metadata_fields(metadata), + ) + return None + + if _oci_upload_config_from_env() is None: + log_flow_event( + active_logger, + "entire_call_recording_disabled", + session_id=session_id, + reason="missing_oci_config", + **_metadata_fields(metadata), + ) + _emit_timeline( + timeline, + "entire_call_recording_disabled", + session_id=session_id, + reason="missing_oci_config", + **_metadata_fields(metadata), + ) + return None + + recorder = EntireCallRecorder( + session_id=session_id, + sample_rate=sample_rate, + channels=channels, + sample_width=sample_width, + frame_ms=frame_ms, + bytes_per_frame=bytes_per_frame, + tmp_dir=_tmp_dir_from_env(), + logger_override=active_logger, + timeline=timeline, + metadata=metadata, + ) + try: + recorder.start() + except Exception: + log_flow_event( + active_logger, + "entire_call_recording_disabled", + session_id=session_id, + reason="start_failed", + **_metadata_fields(metadata), + ) + _emit_timeline( + timeline, + "entire_call_recording_disabled", + session_id=session_id, + reason="start_failed", + **_metadata_fields(metadata), + ) + active_logger.exception( + "[entire-call-recording][start_failed] session_id=%s", + session_id, + ) + return None + return recorder diff --git a/src/app/utils/interruption_tracker.py b/src/app/utils/interruption_tracker.py new file mode 100644 index 0000000..e637b29 --- /dev/null +++ b/src/app/utils/interruption_tracker.py @@ -0,0 +1,418 @@ +# app/utils/interruption_tracker.py +from __future__ import annotations + +import inspect +import logging +import re +import unicodedata +from dataclasses import dataclass +from typing import Any, Callable, Dict, Optional, Set +"""from agent.classifier.interruption_classifier import InterruptionClassifier + +interruption_classifier = InterruptionClassifier() + +def is_backchannel(text: str) -> bool: + result = interruption_classifier.run(text) + print(f"[is_backchannel] {text} -> {result}") + return result""" + +_BACKCHANNEL_SET = { + "ta", "tá", "ok", "okay", "certo", "beleza", "blz", + "tudo bem", "tudo bem?", "tudo bom", "tudo bom?", + "aham", "uhum", "hum", "hm", "uh-huh", + "entendi", "isso", "claro", +} + + +def _norm_pt(s: str) -> str: + s = (s or "").strip().lower() + s = unicodedata.normalize("NFKD", s) + s = "".join(ch for ch in s if not unicodedata.combining(ch)) + s = re.sub(r"\s+", " ", s) + s = re.sub(r"[^\w\s?]", "", s) # remove pontuação (mantém ? e -) + return s.strip() + + +def is_backchannel(text: str) -> bool: + t = _norm_pt(text) + print(f"[is_backchannel] {text} -> {t}") + if not t: + return True + if t in _BACKCHANNEL_SET: + return True + if len(t) <= 2: + return True + #if t.startswith("ta ") or t.startswith("tá "): + # return True + return False + +def remaining_after_prefix(full: str, spoken: str) -> str: + """ + Mais robusto: + - tenta prefixo exato + - tenta common-prefix (char a char) + - tenta find() + - fallback: retorna full (melhor do que ficar em silêncio) + """ + full = full or "" + spoken = spoken or "" + if not full: + return "" + + if not spoken: + return full + + if full.startswith(spoken): + return full[len(spoken):] + + # longest common prefix + m = 0 + for a, b in zip(full, spoken): + if a == b: + m += 1 + else: + break + if m > 0: + return full[m:] + + idx = full.find(spoken) + if idx >= 0: + return full[idx + len(spoken):] + + # fallback: repete o full (melhor do que silêncio) + return full + + +@dataclass +class InterruptionTrace: + stage: str + full_text: str + started_at: float + spoken_text: str = "" + ended_at: float = 0.0 + interrupted: bool = False + interrupter_text: str = "" + interrupter_transcript_raw: str = "" + + +class InterruptionTracker: + """ + Encapsula: + - detecção de backchannel + - rastreio de interrupção (full vs spoken) + - captura do que o usuário disse ao interromper + - auto-resume se for backchannel OU interrupção vazia + - payload "interrupt" para ser passado ao pipeline.run + - interrupt_now() para cortar o áudio IMEDIATAMENTE (sem esperar is_final) + """ + + def __init__( + self, + *, + allow_interruptions_for_stage: Callable[[str], bool], + logger: Optional[logging.Logger] = None, + auto_resume_stages: Optional[Set[str]] = None, + max_auto_resume: int = 2, + auto_resume_on_empty_interrupt: bool = True, + auto_resume_backchannel_any_stage: bool = True, + auto_resume_empty_any_stage: bool = True, + on_interrupt: Optional[Callable[[dict], None]] = None, + ) -> None: + self.allow_interruptions_for_stage = allow_interruptions_for_stage + self.log = logger or logging.getLogger(__name__) + self.on_interrupt = on_interrupt + + self.auto_resume_stages = {s.upper() for s in (auto_resume_stages or {"INTRO", "PRESENTATION"})} + self.max_auto_resume = max(0, int(max_auto_resume)) + self.auto_resume_on_empty_interrupt = bool(auto_resume_on_empty_interrupt) + + self.auto_resume_backchannel_any_stage = bool(auto_resume_backchannel_any_stage) + self.auto_resume_empty_any_stage = bool(auto_resume_empty_any_stage) + + self._loop = None + self._agent_speaking = None + self._session = None + + self.active_trace: Optional[InterruptionTrace] = None + self.last_interrupt_info: Optional[dict] = None + + # ✅ handle do say() atual (para cortar imediatamente) + self._current_handle = None + + # ✅ buffer caso note_user_interrupt chegue antes do active_trace + self._pending_interrupter_text: str = "" + self._pending_interrupter_raw: str = "" + + def is_backchannel_text(self, text: str) -> bool: + return is_backchannel(text) + + def attach(self, session: Any, agent_speaking_event: Any, loop: Any) -> None: + self._loop = loop + self._agent_speaking = agent_speaking_event + self._session = session + session.on("conversation_item_added")(self._on_conversation_item_added) + + def consume_last_interrupt_info(self) -> Optional[dict]: + info = self.last_interrupt_info + self.last_interrupt_info = None + return info + + def note_user_interrupt(self, user_text: str, raw_transcript: str) -> None: + """ + Registra o texto do usuário que interrompeu. + + ✅ Se ainda não existe active_trace (race com callbacks), + guarda num buffer e aplica assim que começar o próximo say_with_trace. + """ + if not user_text: + return + + if self.active_trace and not self.active_trace.interrupter_text: + self.active_trace.interrupter_text = user_text + self.active_trace.interrupter_transcript_raw = raw_transcript or "" + return + + # fallback: guarda para aplicar quando o trace começar + if not self._pending_interrupter_text: + self._pending_interrupter_text = user_text + self._pending_interrupter_raw = raw_transcript or "" + + def should_ignore_user_input(self, user_text: str, agent_is_speaking: bool) -> bool: + if not user_text: + return True + if agent_is_speaking and is_backchannel(user_text): + return True + return False + + def interrupt_now(self) -> bool: + """ + Tenta interromper a fala atual imediatamente. + Retorna True se conseguiu disparar algum método de interrupção. + """ + h = self._current_handle + + # 1) tenta no handle retornado por say() + if h is not None: + for meth in ("interrupt", "cancel", "stop"): + fn = getattr(h, meth, None) + if callable(fn): + try: + ret = fn() + if inspect.isawaitable(ret) and self._loop: + self._loop.create_task(ret) + return True + except Exception: + self.log.exception("interrupt_now falhou usando handle.%s()", meth) + return False + + # 2) fallback: tenta no session (se existir) + s = self._session + if s is not None: + for meth in ("interrupt", "cancel", "stop"): + fn = getattr(s, meth, None) + if callable(fn): + try: + ret = fn() + if inspect.isawaitable(ret) and self._loop: + self._loop.create_task(ret) + return True + except TypeError: + continue + except Exception: + self.log.exception("interrupt_now falhou usando session.%s()", meth) + return False + + return False + + def _on_conversation_item_added(self, ev: Any) -> None: + """ + Captura o item do assistant "commitado". + Quando há interrupção, o texto vem truncado e item.interrupted=True. + """ + if not self.active_trace: + return + + item = getattr(ev, "item", None) + if not item: + return + + role = getattr(item, "role", None) + role_s = str(role).lower() + if "assistant" not in role_s: + return + + try: + self.active_trace.spoken_text = getattr(item, "text_content", "") or "" + self.active_trace.interrupted = bool(getattr(item, "interrupted", False)) + except Exception: + pass + + def force_interrupt_on_end(self, *, reason: str = "call_end") -> Optional[dict]: + """ + Se a call cair enquanto o agente está falando, pode não chegar um item final/truncado. + Cria um payload de interrupção a partir do trace ativo e guarda em last_interrupt_info. + """ + trace = self.active_trace + if not trace: + return None + + # Marca como interrompido por encerramento + trace.interrupted = True + + spoken = trace.spoken_text or "" + full = trace.full_text or "" + remaining = remaining_after_prefix(full, spoken).strip() + + loop = self._loop + now_ms = int(loop.time() * 1000) if loop else 0 + + info = { + "id": f"it-{now_ms}", + "reason": reason, + "stage": trace.stage, + "assistant_full": full, + "assistant_spoken": spoken, + "assistant_remaining": remaining, + "cut_chars": len(spoken), + "ratio_spoken": (len(spoken) / max(1, len(full))) if full else 0.0, + "elapsed_s": 0.0, + "user_interrupt_text": trace.interrupter_text or "", + "user_interrupt_transcript_raw": trace.interrupter_transcript_raw or "", + } + + self.last_interrupt_info = info + + if self.on_interrupt: + try: + self.on_interrupt(info) + except Exception: + self.log.exception("on_interrupt falhou") + + return info + + async def _call_say(self, session: Any, text: str, allow: bool): + res = session.say(text, allow_interruptions=allow, add_to_chat_ctx=True) + if inspect.isawaitable(res): + return await res + return res + + async def say_with_trace(self, session: Any, stage: str, text_to_say: str) -> InterruptionTrace: + stage_u = (stage or "UNKNOWN").upper() + allow = bool(self.allow_interruptions_for_stage(stage_u)) + full = (text_to_say or "").strip() + + if not full: + loop = self._loop + t = loop.time() if loop else 0.0 + return InterruptionTrace(stage=stage_u, full_text="", started_at=t) + + attempt = 0 + last_trace: Optional[InterruptionTrace] = None + + while True: + loop = self._loop + start_t = loop.time() if loop else 0.0 + + trace = InterruptionTrace(stage=stage_u, full_text=full, started_at=start_t) + self.active_trace = trace + + # ✅ aplica interrupter pendente, caso tenha chegado antes do trace existir + if self._pending_interrupter_text and not trace.interrupter_text: + trace.interrupter_text = self._pending_interrupter_text + trace.interrupter_transcript_raw = self._pending_interrupter_raw + self._pending_interrupter_text = "" + self._pending_interrupter_raw = "" + + if self._agent_speaking: + self._agent_speaking.set() + + handle = None + try: + handle = await self._call_say(session, full, allow) + + # ✅ guarda handle atual para interrupt_now() + self._current_handle = handle + + # ✅ aguarda playout + join = getattr(handle, "join", None) + if callable(join): + fut = join() + if fut is not None: + await fut + else: + wait = getattr(handle, "wait_for_playout", None) + if callable(wait): + await wait() + + finally: + end_t = loop.time() if loop else 0.0 + trace.ended_at = end_t + + if self._agent_speaking: + self._agent_speaking.clear() + + self.active_trace = None + self._current_handle = None + + # interrupted flag + trace.interrupted = bool(getattr(handle, "interrupted", False)) if handle else False + + # fallback spoken_text + if trace.interrupted and not trace.spoken_text: + trace.spoken_text = "" + if (not trace.interrupted) and not trace.spoken_text: + trace.spoken_text = trace.full_text + + last_trace = trace + + if trace.interrupted: + remaining = remaining_after_prefix(trace.full_text, trace.spoken_text).strip() + now_ms = int(loop.time() * 1000) if loop else 0 + + info: Dict[str, Any] = { + "id": f"it-{now_ms}", + "reason": "barge_in", + "stage": trace.stage, + "assistant_full": trace.full_text, + "assistant_spoken": trace.spoken_text, + "assistant_remaining": remaining, + "cut_chars": len(trace.spoken_text or ""), + "ratio_spoken": (len(trace.spoken_text) / max(1, len(trace.full_text))) if trace.full_text else 0.0, + "elapsed_s": round(trace.ended_at - trace.started_at, 3), + "user_interrupt_text": trace.interrupter_text, + "user_interrupt_transcript_raw": trace.interrupter_transcript_raw, + } + + self.last_interrupt_info = info + + if self.on_interrupt: + try: + self.on_interrupt(info) + except Exception: + self.log.exception("on_interrupt falhou") + + empty_interrupt = (not (trace.interrupter_text or "").strip()) and self.auto_resume_on_empty_interrupt + back = is_backchannel(trace.interrupter_text) + + stage_ok = stage_u in self.auto_resume_stages + if back and self.auto_resume_backchannel_any_stage: + stage_ok = True + if empty_interrupt and self.auto_resume_empty_any_stage: + stage_ok = True + + can_resume = ( + stage_ok + and attempt < self.max_auto_resume + and (empty_interrupt or back) + ) + + if can_resume and remaining: + attempt += 1 + full = remaining + trace.interrupter_text = "" + trace.interrupter_transcript_raw = "" + continue + + break + + return last_trace or InterruptionTrace(stage=stage_u, full_text=full, started_at=(self._loop.time() if self._loop else 0.0)) diff --git a/src/app/utils/logging.py b/src/app/utils/logging.py new file mode 100644 index 0000000..71cf011 --- /dev/null +++ b/src/app/utils/logging.py @@ -0,0 +1,958 @@ +from __future__ import annotations + +import atexit +import copy +import json +import logging +import os +import queue +import re +import sys +import time +import uuid +from collections.abc import Iterable, Mapping +from contextvars import ContextVar +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from threading import Event, Lock, Thread +from typing import Any + +from app.utils.structured_otlp import publish_structured_span +from app.utils.structured_pubsub import publish_structured_event + +_LOGGER_CACHE: dict[str, logging.Logger] = {} +_LOGGER_LOCK = Lock() +_BASE_LOGGER_INITIALIZED = False +_LOG_SESSION_ID: ContextVar[str] = ContextVar("log_session_id", default="") + + +def set_log_session_id(session_id: str | None) -> None: + """Bind the call session to the current async context and its child tasks.""" + _LOG_SESSION_ID.set(str(session_id or "").strip()) + + +class _SessionConsoleFormatter(logging.Formatter): + def __init__(self, fmt: str, *, fixed_session_id: str = "") -> None: + super().__init__(fmt) + self._fixed_session_id = str(fixed_session_id or "").strip() + + def format(self, record: logging.LogRecord) -> str: + rendered = super().format(record) + session_id = str( + getattr(record, "session_id", "") + or self._fixed_session_id + or _LOG_SESSION_ID.get() + or "" + ).strip() + if not session_id or re.search( + r"(?i)(?:[\"']session_?id[\"']|\bsession_?id\b)\s*[:=]", + rendered, + ): + return rendered + return f"{rendered} | session_id={session_id}" + + +class _CallFileFormatter(logging.Formatter): + def format(self, record: logging.LogRecord) -> str: + match getattr(record, "structured_event", False): + case True: + return record.getMessage() + case _: + return super().format(record) + + +def _resolve_level(default: int) -> int: + raw = (os.getenv("LOG_LEVEL", "") or "").strip().upper() + match raw: + case "": + return default + case _: + value = logging.getLevelName(raw) + match value: + case int(): + return value + case _: + return default + + +EVENT_RECEBIMENTO_MSG = "recebimento msg" +EVENT_ENVIO_MSG = "envio msg" + +_FINALIZATION_MATCHERS: tuple[tuple[str, tuple[str, ...]], ...] = ( + ("falha_stt", ("stt",)), + ("falha_tts", ("tts",)), + ("falha_agente", ("agent_backend", "falha agente")), + ("falha_tia", ("agent_runtime", "bridge", "capacity")), + ("silencio_longo", ("no_user_response", "silencio", "sil\u00eancio")), + ("outro_assunto", ("outro_assunto", "other_subject")), + ("nao_resolvido", ("nao_resolvido", "n\u00e3o resolvido", "unresolved")), + ("resolvido", ("stage_done", "resolved", "resolvido", "done")), + ("hang_up", ("disconnect", "hang")), +) + +_ERROR_MATCHERS: tuple[tuple[str, tuple[str, ...]], ...] = ( + ("transferido", ("transferred", "transferido", "transferencia", "transfer\u00eancia")), + ("silencio_longo", ("stop_silencio_longo", "no_user_response", "silencio", "sil\u00eancio")), + ("capacidade_tia", ("capacity",)), + ("falha_tia", ("agent_runtime", "livekit")), + ("falha_stt", ("stt",)), + ("falha_tts", ("tts",)), + ("falha_comunicacao_agente", ("agent_backend", "remote_agent", "falha agente")), + ("falha_comunicacao_gstreamer", ("bridge", "gstreamer", "gst")), + ("capacidade_agente", ("agent_runtime",)), +) + + +@dataclass(frozen=True, slots=True) +class StructuredLogContext: + callid: str + session_id: str + num_telefone: str + cod_ani: str + nome_agente: str + message_id_seed: str = "" + + def as_metadata(self) -> dict[str, str]: + return { + "callid": self.callid, + "session_id": self.session_id, + "num_telefone": self.num_telefone, + "cod_ani": self.cod_ani, + "nome_agente": self.nome_agente, + "message_id_seed": self.message_id_seed, + } + + +def _pick_value(payload: Mapping[str, Any] | None, *keys: str) -> str: + match payload: + case Mapping(): + pass + case _: + return "" + + for key in keys: + value = payload.get(key) + if value not in (None, ""): + return str(value).strip() + + lowered = {str(key).lower(): value for key, value in payload.items()} + for key in keys: + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return str(value).strip() + + return "" + + +def normalize_agent_log_name(value: Any) -> str: + raw = str(value or "").strip().lower() + match raw: + case "conta" | "contas": + return "conta" + case "ofert" | "oferta" | "ofertas": + return "ofert" + case "cobra" | "cobranca" | "cobran\u00e7a" | "cobrancas" | "cobran\u00e7as": + return "cobra" + case "": + return "unknown" + case _: + return raw + + +def build_callid( + *, + router_call_key_day: str = "", + router_call_key: str = "", + call_id_ged: str = "", + protocol: str = "", + room: str = "", +) -> str: + day = str(router_call_key_day or "").strip() + key = str(router_call_key or "").strip() + match (day, key): + case ("", ""): + return str(call_id_ged or protocol or room or "unknown").strip() or "unknown" + case _: + return f"{day}{key}" + + +def structured_context_from_start_data( + data: Mapping[str, Any] | None, + *, + session_id: str = "", +) -> StructuredLogContext: + router_day = _pick_value(data, "routerCallKeyDay", "RouterCallKeyDay", "router_call_key_day") + router_key = _pick_value(data, "routerCallKey", "RouterCallKey", "router_call_key") + call_id_ged = _pick_value(data, "callIdGed", "call_id_ged") + callid = build_callid( + router_call_key_day=router_day, + router_call_key=router_key, + call_id_ged=call_id_ged, + ) + return StructuredLogContext( + callid=callid, + session_id=( + str(session_id or "").strip() + or _pick_value(data, "session_id", "sessionId") + ), + num_telefone=_pick_value(data, "gsm", "GSM", "msisdn", "NUM_TELEFONE", "phone"), + cod_ani=_pick_value(data, "ani", "ANI", "cod_ani"), + nome_agente=normalize_agent_log_name(_pick_value(data, "agent", "agente", "nome_agente", "Nome_agente")), + message_id_seed=call_id_ged or f"MSG-{callid}", + ) + + +def structured_context_from_metadata(metadata: Mapping[str, Any] | None) -> StructuredLogContext: + match metadata: + case Mapping(): + structured_context = metadata.get("structured_log_context") + case _: + structured_context = None + + match structured_context: + case Mapping(): + return StructuredLogContext( + callid=_pick_value(structured_context, "callid") or "unknown", + session_id=_pick_value(structured_context, "session_id", "sessionId"), + num_telefone=_pick_value(structured_context, "num_telefone") or "unknown", + cod_ani=_pick_value(structured_context, "cod_ani") or "unknown", + nome_agente=normalize_agent_log_name(_pick_value(structured_context, "nome_agente", "Nome_agente")), + message_id_seed=_pick_value(structured_context, "message_id_seed"), + ) + case _: + pass + + match metadata: + case Mapping(): + session_data = metadata.get("session_data") + remote_agent = metadata.get("remote_agent") + case _: + session_data = {} + remote_agent = {} + + callid = build_callid( + router_call_key_day=_pick_value(remote_agent, "RouterCallKeyDay", "routerCallKeyDay"), + router_call_key=_pick_value(remote_agent, "RouterCallKey", "routerCallKey"), + call_id_ged=_pick_value(metadata, "call_id_ged", "callIdGed"), + protocol=_pick_value(metadata, "protocol"), + room=_pick_value(metadata, "timeline_id"), + ) + + return StructuredLogContext( + callid=callid, + session_id=( + _pick_value(metadata, "session_id", "sessionId") + or _pick_value(remote_agent, "session_id", "sessionId") + ), + num_telefone=( + _pick_value(session_data, "gsm", "GSM", "msisdn", "NUM_TELEFONE", "phone") + or _pick_value(remote_agent, "GSM", "gsm", "NUM_TELEFONE") + or "unknown" + ), + cod_ani=_pick_value(session_data, "ani", "ANI") or _pick_value(remote_agent, "ANI", "ani") or "unknown", + nome_agente=normalize_agent_log_name( + _pick_value(remote_agent, "agent", "agente") or _pick_value(session_data, "agent", "agente") + ), + message_id_seed=( + _pick_value(metadata, "call_id_ged", "callIdGed") + or _pick_value(remote_agent, "callIdGed", "call_id_ged") + or f"MSG-{callid}" + ), + ) + + +def format_event_timestamp(ns: int | None = None) -> str: + match ns: + case None: + ts_ns = time.time_ns() + case _: + ts_ns = int(ns) + seconds, remainder_ns = divmod(ts_ns, 1_000_000_000) + dt = datetime.fromtimestamp(seconds, tz=timezone.utc) + return f"{dt:%Y-%m-%dT%H:%M:%S}.{remainder_ns // 1_000_000:03d}Z" + + +def _format_optional_event_timestamp(ns: int | None) -> str | None: + match ns: + case None: + return None + case _: + return format_event_timestamp(ns) + + +def event_latency_ms(start_ns: int | None, end_ns: int | None = None) -> int | None: + match start_ns: + case None: + return None + case _: + started_ns = int(start_ns) + + match end_ns: + case None: + finished_ns = time.time_ns() + case _: + finished_ns = int(end_ns) + return max(0, round((finished_ns - started_ns) / 1_000_000)) + + +def _join_log_parts(*parts: str) -> str: + return " ".join(str(item or "").strip().lower() for item in parts if str(item or "").strip()) + + +def _match_first_token(text: str, matchers: tuple[tuple[str, tuple[str, ...]], ...]) -> str: + for key, tokens in matchers: + if any(token in text for token in tokens): + return key + return "" + + +def _structured_event_kind(tipo_evento: str = "") -> str: + normalized = str(tipo_evento or "").strip().lower().replace("_", " ") + match normalized: + case "recebimento" | "recebimento msg": + return "recebimento" + case "envio" | "envio msg": + return "envio" + case _: + return "" + + +def _event_kind_allows(kind: str, *allowed: str) -> bool: + return not kind or kind in allowed + + +def finalization_from_status( + *, + status: str = "", + reason: str = "", + resource: str = "", +) -> str | None: + normalized = _join_log_parts(status, reason, resource) + match normalized: + case "": + return None + case _: + pass + + match _match_first_token(normalized, _FINALIZATION_MATCHERS): + case "falha_stt": + return "Falha STT" + case "falha_tts": + return "Falha TTS" + case "falha_agente": + return "Falha Agente" + case "falha_tia": + return "Falha TIA" + case "silencio_longo": + return "Sil\u00eancio Longo" + case "outro_assunto": + return "Outro Assunto" + case "nao_resolvido": + return "N\u00e3o Resolvido" + case "resolvido": + return "Resolvido" + case "hang_up": + return "hang up" + case _: + return reason or status or None + + +def error_message_from_resource( + resource: str = "", + status: str = "", + reason: str = "", + *, + tipo_evento: str = "", +) -> str | None: + normalized = _join_log_parts(resource, status, reason) + match normalized: + case "": + return None + case _: + pass + + event_kind = _structured_event_kind(tipo_evento) + match _match_first_token(normalized, _ERROR_MATCHERS): + case "transferido" if _event_kind_allows(event_kind, "envio"): + return "Transferido" + case "silencio_longo" if _event_kind_allows(event_kind, "recebimento"): + return "Silencio Longo" + case "capacidade_tia": + return "Capacidade TIA" + case "falha_tia" if _event_kind_allows(event_kind, "recebimento", "envio"): + return "Falha TIA" + case "falha_stt" if _event_kind_allows(event_kind, "recebimento"): + return "Falha STT" + case "falha_tts" if _event_kind_allows(event_kind, "envio"): + return "Falha TTS" + case "falha_comunicacao_agente" if _event_kind_allows(event_kind, "envio"): + return "Falha comunicacao" + case "falha_comunicacao_gstreamer" if _event_kind_allows(event_kind, "recebimento"): + return "Falha comunicacao" + case "capacidade_agente": + return "Capacidade Agente" + case _: + return None + + +def _event_field(value: Any) -> Any: + return "" if value is None else value + + +def _event_flag(value: Any) -> int | None: + if value is None: + return None + if isinstance(value, str): + normalized = value.strip().lower() + if not normalized: + return None + if normalized in {"1", "true", "yes", "on"}: + return 1 + if normalized in {"0", "false", "no", "off"}: + return 0 + return 1 if bool(value) else 0 + + +def _env_flag(name: str, default: str = "1") -> bool: + return (os.getenv(name, default) or default).strip().lower() in {"1", "true", "yes", "on"} + + +def _flow_value(value: Any) -> str: + text = str(value) + text = " ".join(text.split()) + return text.replace("|", "/") + + +def _flow_preview(value: Any, *, max_chars: int | None = None) -> str: + if max_chars is None: + try: + max_chars = int(os.getenv("FLOW_LOG_PREVIEW_CHARS", "500")) + except ValueError: + max_chars = 500 + text = _flow_value(value) + if len(text) <= max_chars: + return text + return f"{text[:max_chars]}..." + + +def log_flow_event(logger: Any, step: str, **fields: Any) -> None: + if not _env_flag("FLOW_LOG_ENABLED", "1"): + return + + safe_fields: dict[str, str] = {} + for key, value in fields.items(): + if value is None: + continue + if isinstance(value, str) and not value.strip(): + continue + if key in {"text", "transcript", "reply", "payload", "result", "output", "content"}: + safe_fields[key] = _flow_preview(value) + else: + safe_fields[key] = _flow_value(value) + + parts = [f"step={_flow_value(step or 'unknown')}"] + parts.extend(f"{key}={value}" for key, value in safe_fields.items()) + logger.info("FLOW | %s", " | ".join(parts)) + + +def build_structured_event( + context: StructuredLogContext, + *, + tipo_evento: str, + message_id: str | None = None, + inicio_ns: int | None = None, + fim_ns: int | None = None, + latencia_total_ms: int | None = None, + latencia_tffb_ms: int | None = None, + duracao_audio_ms: int | None = None, + tts_max_gap_ms: int | None = None, + tts_max_underrun_0ms: int | None = None, + tts_underflow_count: int | None = None, + tts_avg_underflow_ms: int | None = None, + interrupcao: bool | int | str | None = None, + erro_msg: str | None = None, + erro_detalhe: str | None = None, + http_cod_status: int | str | None = None, + http_cod_desc: str | None = None, + finalizacao: str | None = None, +) -> dict[str, Any]: + return { + "tipo_evento": str(tipo_evento or ""), + "message_id": str(message_id or uuid.uuid4()), + "dat_hora_inicio": format_event_timestamp(inicio_ns), + "dat_hora_fim": _event_field(_format_optional_event_timestamp(fim_ns)), + "callid": str(context.callid or ""), + "session_id": str(context.session_id or ""), + "num_telefone": str(context.num_telefone or ""), + "cod_ani": str(context.cod_ani or ""), + "latencia_total_STT_TTS": _event_field(latencia_total_ms), + "latencia_TFFB_STT_TTS": _event_field(latencia_tffb_ms), + "duracao_audio": _event_field(duracao_audio_ms), + "tts_max_gap_ms": _event_field(tts_max_gap_ms), + "tts_max_underrun_0ms": _event_field(tts_max_underrun_0ms), + "tts_underflow_count": _event_field(tts_underflow_count), + "tts_avg_underflow_ms": _event_field(tts_avg_underflow_ms), + "interrupcao": _event_field(_event_flag(interrupcao)), + "erro_msg": _event_field(erro_msg), + "erro_detalhe": _event_field(erro_detalhe), + "http_cod_status": _event_field(http_cod_status), + "http_cod_desc": _event_field(http_cod_desc), + "finalizacao": _event_field(finalizacao), + "nome_agente": str(context.nome_agente or ""), + } + + +def log_structured_event( + logger: Any, + context: StructuredLogContext | None, + *, + tipo_evento: str, + message_id: str | None = None, + inicio_ns: int | None = None, + fim_ns: int | None = None, + latencia_total_ms: int | None = None, + latencia_tffb_ms: int | None = None, + duracao_audio_ms: int | None = None, + tts_max_gap_ms: int | None = None, + tts_max_underrun_0ms: int | None = None, + tts_underflow_count: int | None = None, + tts_avg_underflow_ms: int | None = None, + interrupcao: bool | int | str | None = None, + erro_msg: str | None = None, + erro_detalhe: str | None = None, + http_cod_status: int | str | None = None, + http_cod_desc: str | None = None, + finalizacao: str | None = None, +) -> dict[str, Any] | None: + match (os.getenv("STRUCTURED_EVENT_LOG_ENABLED", "1"), context): + case ("1", StructuredLogContext()): + pass + case _: + return None + + event = build_structured_event( + context, + tipo_evento=tipo_evento, + message_id=message_id, + inicio_ns=inicio_ns, + fim_ns=fim_ns, + latencia_total_ms=latencia_total_ms, + latencia_tffb_ms=latencia_tffb_ms, + duracao_audio_ms=duracao_audio_ms, + tts_max_gap_ms=tts_max_gap_ms, + tts_max_underrun_0ms=tts_max_underrun_0ms, + tts_underflow_count=tts_underflow_count, + tts_avg_underflow_ms=tts_avg_underflow_ms, + interrupcao=interrupcao, + erro_msg=erro_msg, + erro_detalhe=erro_detalhe, + http_cod_status=http_cod_status, + http_cod_desc=http_cod_desc, + finalizacao=finalizacao, + ) + logger.info(json.dumps(event, ensure_ascii=False), extra={"structured_event": True}) + publish_structured_event(event) + publish_structured_span(event) + return event + + +def _sanitize_for_filename(value: str) -> str: + value = (value or "").strip() + match value: + case "": + return "unknown" + case _: + pass + value = re.sub(r"\D+", "", value) + return value or "unknown" + + +def _bounded_env_int(name: str, default: int, *, minimum: int, maximum: int) -> int: + try: + value = int(os.getenv(name, str(default)) or str(default)) + except (TypeError, ValueError): + value = default + return max(minimum, min(maximum, value)) + + +class _LogWriterCommand: + __slots__ = ("kind", "event") + + def __init__(self, kind: str) -> None: + self.kind = kind + self.event = Event() + + +class _LogWriter: + """Executa handlers de arquivo numa thread dedicada e observável.""" + + def __init__( + self, + max_queue: int, + *, + warning_interval_s: float = 60.0, + ) -> None: + self._queue: queue.Queue[ + tuple[logging.Handler, logging.LogRecord] | _LogWriterCommand + ] = queue.Queue(maxsize=max(1, max_queue)) + self._thread: Thread | None = None + self._lifecycle_lock = Lock() + self._stats_lock = Lock() + self._closed = False + self._dropped = 0 + self._write_errors = 0 + self._last_drop_warning_at = 0.0 + self._last_error_warning_at = 0.0 + self._warning_interval_s = max(1.0, float(warning_interval_s)) + self._logger = logging.getLogger(f"{__name__}.writer") + + @property + def dropped(self) -> int: + with self._stats_lock: + return self._dropped + + @property + def write_errors(self) -> int: + with self._stats_lock: + return self._write_errors + + def start(self) -> bool: + with self._lifecycle_lock: + if self._closed: + return False + if self._thread is not None and self._thread.is_alive(): + return True + self._thread = Thread( + target=self._run, + name="call-log-writer", + daemon=True, + ) + self._thread.start() + return True + + def submit(self, target: logging.Handler, record: logging.LogRecord) -> bool: + if not self.start(): + self._record_drop(reason="writer_closed") + return False + reason = "" + with self._lifecycle_lock: + if self._closed: + reason = "writer_closed" + else: + try: + self._queue.put_nowait((target, record)) + return True + except queue.Full: + reason = "queue_full" + self._record_drop(reason=reason) + return False + + def flush(self, timeout: float | None = 2.0) -> bool: + with self._lifecycle_lock: + if self._closed: + return self._queue.empty() + if not self.start(): + return False + command = _LogWriterCommand("flush") + deadline = None if timeout is None else time.monotonic() + max(0.0, timeout) + with self._lifecycle_lock: + if self._closed: + return self._queue.empty() + if not self._enqueue_command(command, deadline): + return False + remaining = None if deadline is None else max(0.0, deadline - time.monotonic()) + return command.event.wait(remaining) + + def shutdown(self, timeout: float | None = 2.0) -> bool: + with self._lifecycle_lock: + if self._closed: + thread = self._thread + return thread is None or not thread.is_alive() + self._closed = True + thread = self._thread + + if thread is None: + return True + + command = _LogWriterCommand("shutdown") + deadline = None if timeout is None else time.monotonic() + max(0.0, timeout) + if not self._enqueue_command(command, deadline): + return False + remaining = None if deadline is None else max(0.0, deadline - time.monotonic()) + if not command.event.wait(remaining): + return False + remaining = None if deadline is None else max(0.0, deadline - time.monotonic()) + thread.join(remaining) + return not thread.is_alive() + + def _enqueue_command( + self, + command: _LogWriterCommand, + deadline: float | None, + ) -> bool: + try: + if deadline is None: + self._queue.put(command) + else: + self._queue.put(command, timeout=max(0.0, deadline - time.monotonic())) + return True + except queue.Full: + return False + + def _run(self) -> None: + while True: + item = self._queue.get() + if isinstance(item, _LogWriterCommand): + item.event.set() + if item.kind == "shutdown": + return + continue + + target, record = item + try: + target.handle(record) + except Exception as exc: + self._record_write_error(target=target, exc=exc) + + def _record_drop(self, *, reason: str) -> None: + now = time.monotonic() + with self._stats_lock: + self._dropped += 1 + dropped = self._dropped + should_warn = now - self._last_drop_warning_at >= self._warning_interval_s + if should_warn: + self._last_drop_warning_at = now + if should_warn: + self._logger.warning( + "CALL_LOG_ASYNC_DROP | reason=%s | dropped_total=%s | queue_size=%s", + reason, + dropped, + self._queue.qsize(), + ) + + def _record_write_error( + self, + *, + target: logging.Handler, + exc: Exception, + ) -> None: + now = time.monotonic() + with self._stats_lock: + self._write_errors += 1 + errors = self._write_errors + should_warn = now - self._last_error_warning_at >= self._warning_interval_s + if should_warn: + self._last_error_warning_at = now + if should_warn: + self._logger.warning( + "CALL_LOG_ASYNC_WRITE_FAIL | handler=%s | error=%s: %s | errors_total=%s", + type(target).__name__, + type(exc).__name__, + exc, + errors, + ) + + +_LOG_WRITER = _LogWriter( + _bounded_env_int( + "CALL_LOG_QUEUE_MAX", + 20_000, + minimum=1, + maximum=1_000_000, + ), + warning_interval_s=_bounded_env_int( + "ASYNC_IO_WARNING_INTERVAL_S", + 60, + minimum=1, + maximum=3_600, + ), +) +atexit.register(_LOG_WRITER.shutdown) + + +class _AsyncFileHandler(logging.Handler): + """Enfileira uma cópia estável do registro para escrita em background.""" + + def __init__(self, target: logging.Handler) -> None: + super().__init__(level=target.level) + self._target = target + _LOG_WRITER.start() + + def emit(self, record: logging.LogRecord) -> None: + clone = copy.copy(record) + try: + clone.msg = clone.getMessage() + clone.args = None + except Exception: + pass + _LOG_WRITER.submit(self._target, clone) + + def flush(self) -> None: + writer = globals().get("_LOG_WRITER") + if writer is not None: + writer.flush(timeout=2.0) + + def close(self) -> None: + try: + self.flush() + finally: + try: + self._target.close() + finally: + super().close() + + +def _build_call_file_handler(file_path: Path, level: int) -> logging.FileHandler: + handler = logging.FileHandler(file_path, encoding="utf-8") + handler.setLevel(level) + handler.setFormatter( + _CallFileFormatter( + "%(asctime)s | %(levelname)s | %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", + ) + ) + return handler + + +def setup_minimal_logging(*, level: int = logging.INFO) -> logging.Logger: + """ + - Root: WARNING + - Logger base: console simples + - Não cria arquivo por processo + - Arquivo apenas por chamada + """ + global _BASE_LOGGER_INITIALIZED + + level = _resolve_level(level) + logger_name = os.getenv("APP_LOGGER_NAME", "agent_internal_stt") + + match _BASE_LOGGER_INITIALIZED: + case True: + log = logging.getLogger(logger_name) + log.setLevel(level) + return log + case _: + pass + + root = logging.getLogger() + root.setLevel(logging.WARNING) + + for h in list(root.handlers): + root.removeHandler(h) + try: + h.close() + except Exception: + pass + + root_handler = logging.StreamHandler(sys.stdout) + root_handler.setLevel(logging.WARNING) + root_handler.setFormatter( + _SessionConsoleFormatter("%(levelname)s %(name)s: %(message)s") + ) + root.addHandler(root_handler) + + noisy: Iterable[str] = ( + "livekit", + "livekit.agents", + "livekit.plugins", + "livekit.rtc", + "aioice", + "aiortc", + "httpx", + "httpcore", + "asyncio", + "stt_internal", + "transformers", + "uvicorn.access", + ) + for name in noisy: + logging.getLogger(name).setLevel(logging.WARNING) + + log = logging.getLogger(logger_name) + log.setLevel(level) + log.propagate = False + + for h in list(log.handlers): + log.removeHandler(h) + try: + h.close() + except Exception: + pass + + console_handler = logging.StreamHandler(sys.stdout) + console_handler.setLevel(level) + console_handler.setFormatter(_SessionConsoleFormatter("%(message)s")) + log.addHandler(console_handler) + + _BASE_LOGGER_INITIALIZED = True + return log + + +def get_call_logger( + *, + phone_number: str | None = None, + session_id: str | None = None, + level: int = logging.INFO, +) -> logging.Logger: + """ + Logger dedicado por ligação. + Nome do arquivo: telefone_timestamp.txt + Ex: 68981100139_20260306_142317.txt + """ + log_dir = Path(os.getenv("LOG_DIR", "./logs")) + logger_name = os.getenv("APP_LOGGER_NAME", "agent_internal_stt") + log_to_file = os.getenv("LOG_TO_FILE", "1") == "1" + level = _resolve_level(level) + + safe_phone = _sanitize_for_filename(phone_number or "") + safe_session = _sanitize_for_filename(session_id or "") + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + + cache_key = f"{logger_name}.call.{safe_phone}.{safe_session or timestamp}" + + with _LOGGER_LOCK: + cached = _LOGGER_CACHE.get(cache_key) + match cached: + case logging.Logger(): + return cached + case _: + pass + + log = logging.getLogger(cache_key) + log.setLevel(level) + log.propagate = False + + for h in list(log.handlers): + log.removeHandler(h) + try: + h.close() + except Exception: + pass + + console_handler = logging.StreamHandler(sys.stdout) + console_handler.setLevel(level) + console_handler.setFormatter( + _SessionConsoleFormatter("%(message)s", fixed_session_id=safe_session) + ) + log.addHandler(console_handler) + + match log_to_file: + case True: + log_dir.mkdir(parents=True, exist_ok=True) + + file_name = f"{safe_phone}_{timestamp}.txt" + file_path = log_dir / file_name + + log.addHandler(_AsyncFileHandler(_build_call_file_handler(file_path, level))) + log.info( + "CALL_FILE_LOG_ENABLED | phone=%s | session_id=%s | path=%s", + safe_phone, + safe_session or "-", + file_path, + ) + case _: + pass + + _LOGGER_CACHE[cache_key] = log + return log diff --git a/src/app/utils/pcm.py b/src/app/utils/pcm.py new file mode 100644 index 0000000..62e602c --- /dev/null +++ b/src/app/utils/pcm.py @@ -0,0 +1,42 @@ +import math +import audioop +from array import array + +def pcm_duration_ms(pcm: bytes, sample_rate: int, channels: int) -> float: + # PCM16LE => 2 bytes por sample por canal + if not pcm: + return 0.0 + samples_total = len(pcm) // 2 # total de int16 (inclui canais) + samples_per_channel = samples_total / max(1, channels) + return (samples_per_channel / sample_rate) * 1000.0 + +def dbfs_pcm16le(pcm: bytes) -> float: + """ + dBFS aproximado do PCM16LE. + 0 dBFS = full scale, silêncio tende a -inf (aqui clamp em -120). + """ + if not pcm: + return -120.0 + rms = audioop.rms(pcm, 2) # width=2 bytes (int16) + if rms <= 0: + return -120.0 + return 20.0 * math.log10(rms / 32768.0) + +def pcm16_dbfs(frame_bytes: bytes) -> float: + # frame_bytes: PCM16LE mono + if not frame_bytes: + return -120.0 + samples = array("h") + samples.frombytes(frame_bytes) + if len(samples) == 0: + return -120.0 + s2 = 0.0 + for x in samples: + s2 += float(x) * float(x) + rms = math.sqrt(s2 / len(samples)) if s2 > 0 else 0.0 + if rms <= 1e-9: + return -120.0 + return 20.0 * math.log10(rms / 32768.0) + +def is_silence(frame_bytes: bytes, *, threshold_dbfs: float = -50.0) -> bool: + return pcm16_dbfs(frame_bytes) < threshold_dbfs \ No newline at end of file diff --git a/src/app/utils/structured_otlp.py b/src/app/utils/structured_otlp.py new file mode 100644 index 0000000..761266a --- /dev/null +++ b/src/app/utils/structured_otlp.py @@ -0,0 +1,240 @@ +from __future__ import annotations + +import atexit +import json +import os +import re +import sys +import time +from collections.abc import Mapping +from contextvars import ContextVar +from datetime import datetime, timezone +from threading import Lock, Thread +from typing import Any + +from opentelemetry.context import Context +from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter +from opentelemetry.sdk.resources import SERVICE_NAME, Resource +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import BatchSpanProcessor +from opentelemetry.sdk.trace.id_generator import IdGenerator, RandomIdGenerator +from opentelemetry.trace import Status, StatusCode + +_TRACER_PROVIDER: Any | None = None +_TRACER: Any | None = None +_TRACER_ENDPOINT = "" +_TRACER_LOCK = Lock() +_CURRENT_TRACE_ID: ContextVar[int | None] = ContextVar("structured_otlp_trace_id", default=None) + +_HEX_RE = re.compile(r"[^0-9a-fA-F]+") +_ISO_EVENT_TIMESTAMP_RE = re.compile( + r"^(\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2})(?:\.(\d{1,9}))?Z$" +) +_LEGACY_EVENT_TIMESTAMP_RE = re.compile(r"^(\d{2}/\d{2}/\d{4} \d{2}:\d{2}:\d{2}),(\d{1,9})$") +_EXCLUDED_ATTRIBUTE_KEYS = {"original_message_id"} +_OMIT_EMPTY_ATTRIBUTE_KEYS = {"latencia_TFFB_STT_TTS", "duracao_audio", "interrupcao"} + + +class _SessionTraceIdGenerator(IdGenerator): + def __init__(self) -> None: + self._fallback = RandomIdGenerator() + + def generate_span_id(self) -> int: + return self._fallback.generate_span_id() + + def generate_trace_id(self) -> int: + trace_id = _CURRENT_TRACE_ID.get() + if trace_id is not None: + return trace_id + return self._fallback.generate_trace_id() + + +_ID_GENERATOR = _SessionTraceIdGenerator() + + +def trace_id_from_session_id(session_id: Any) -> str: + trace_id = _HEX_RE.sub("", str(session_id or "")).lower() + if len(trace_id) != 32 or trace_id == ("0" * 32): + return "" + return trace_id + + +def _configured_endpoint() -> str: + return (os.getenv("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT", "") or "").strip() + + +def _service_name() -> str: + return ( + os.getenv("OTEL_SERVICE_NAME", "") + or os.getenv("APP_LOGGER_NAME", "") + or "agent_internal_stt" + ).strip() + + +def shutdown_structured_otlp() -> None: + """Drena e encerra o provider corrente no shutdown do processo.""" + global _TRACER_PROVIDER, _TRACER, _TRACER_ENDPOINT + + with _TRACER_LOCK: + provider = _TRACER_PROVIDER + _TRACER_PROVIDER = None + _TRACER = None + _TRACER_ENDPOINT = "" + + if provider is not None: + try: + provider.shutdown() + except Exception as exc: + print( + f"STRUCTURED_EVENT_OTLP_SHUTDOWN_FAIL | error={type(exc).__name__}: {exc}", + file=sys.stderr, + ) + + +def force_flush_structured_otlp(timeout_millis: int = 5_000) -> bool: + """Força o envio pendente fora do caminho de áudio, útil em teardown/testes.""" + with _TRACER_LOCK: + provider = _TRACER_PROVIDER + if provider is None: + return True + try: + return bool(provider.force_flush(timeout_millis=max(0, int(timeout_millis)))) + except Exception as exc: + print( + f"STRUCTURED_EVENT_OTLP_FLUSH_FAIL | error={type(exc).__name__}: {exc}", + file=sys.stderr, + ) + return False + + +def _shutdown_replaced_provider(provider: Any) -> None: + try: + provider.shutdown() + except Exception: + pass + + +def _tracer() -> Any | None: + global _TRACER_PROVIDER, _TRACER, _TRACER_ENDPOINT + + endpoint = _configured_endpoint() + if not endpoint: + return None + + with _TRACER_LOCK: + if _TRACER is not None and _TRACER_ENDPOINT == endpoint: + return _TRACER + + try: + resource = Resource.create(attributes={SERVICE_NAME: _service_name()}) + provider = TracerProvider(resource=resource, id_generator=_ID_GENERATOR) + provider.add_span_processor( + BatchSpanProcessor(OTLPSpanExporter(endpoint=endpoint)) + ) + + previous_provider = _TRACER_PROVIDER + _TRACER_PROVIDER = provider + _TRACER = provider.get_tracer("app.utils.structured_otlp") + _TRACER_ENDPOINT = endpoint + tracer = _TRACER + except Exception as exc: + print( + f"STRUCTURED_EVENT_OTLP_DISABLED | error={type(exc).__name__}: {exc}", + file=sys.stderr, + ) + return None + + if previous_provider is not None: + Thread( + target=_shutdown_replaced_provider, + args=(previous_provider,), + name="structured-otlp-old-provider-shutdown", + daemon=True, + ).start() + return tracer + + +atexit.register(shutdown_structured_otlp) + + +def _parse_event_timestamp(value: Any) -> int | None: + raw = str(value or "").strip() + iso_match = _ISO_EVENT_TIMESTAMP_RE.match(raw) + if iso_match is not None: + dt = datetime.strptime(iso_match.group(1), "%Y-%m-%dT%H:%M:%S").replace(tzinfo=timezone.utc) + fractional_ns = int((iso_match.group(2) or "").ljust(9, "0")[:9]) + return int(dt.timestamp()) * 1_000_000_000 + fractional_ns + + legacy_match = _LEGACY_EVENT_TIMESTAMP_RE.match(raw) + if legacy_match is None: + return None + + dt = datetime.strptime(legacy_match.group(1), "%d/%m/%Y %H:%M:%S") + fractional_ns = int(legacy_match.group(2).ljust(9, "0")[:9]) + return int(dt.timestamp()) * 1_000_000_000 + fractional_ns + + +def _event_times(event: Mapping[str, Any]) -> tuple[int, int]: + start_ns = _parse_event_timestamp(event.get("dat_hora_inicio")) or time.time_ns() + end_ns = _parse_event_timestamp(event.get("dat_hora_fim")) or start_ns + if end_ns < start_ns: + return start_ns, start_ns + return start_ns, end_ns + + +def _attribute_value(value: Any) -> str | bool | int | float | None: + if value is None: + return None + if isinstance(value, str | bool | int | float): + return value + try: + return json.dumps(value, ensure_ascii=False) + except Exception: + return str(value) + + +def _attributes(event: Mapping[str, Any]) -> dict[str, str | bool | int | float]: + attrs: dict[str, str | bool | int | float] = {} + for key, value in event.items(): + if str(key) in _EXCLUDED_ATTRIBUTE_KEYS: + continue + if str(key) in _OMIT_EMPTY_ATTRIBUTE_KEYS and value == "": + continue + attr_value = _attribute_value(value) + if attr_value is not None: + attrs[str(key)] = attr_value + return attrs + + +def publish_structured_span(event: Mapping[str, Any]) -> None: + trace_id = trace_id_from_session_id(event.get("session_id")) + if not trace_id: + return + + tracer = _tracer() + if tracer is None: + return + + try: + start_ns, end_ns = _event_times(event) + tipo_evento = str(event.get("tipo_evento") or "unknown").strip() or "unknown" + + token = _CURRENT_TRACE_ID.set(int(trace_id, 16)) + try: + span = tracer.start_span( + f"structured_log.{tipo_evento}", + context=Context(), + attributes=_attributes(event), + start_time=start_ns, + ) + finally: + _CURRENT_TRACE_ID.reset(token) + + try: + erro_msg = str(event.get("erro_msg") or "").strip() + if erro_msg: + span.set_status(Status(StatusCode.ERROR, erro_msg)) + finally: + span.end(end_time=end_ns) + except Exception as exc: + print(f"STRUCTURED_EVENT_OTLP_PUBLISH_FAIL | error={type(exc).__name__}: {exc}", file=sys.stderr) diff --git a/src/app/utils/structured_pubsub.py b/src/app/utils/structured_pubsub.py new file mode 100644 index 0000000..8f3c771 --- /dev/null +++ b/src/app/utils/structured_pubsub.py @@ -0,0 +1,128 @@ +from __future__ import annotations + +import json +import os +import re +import sys +from collections.abc import Mapping +from datetime import datetime, timedelta, timezone +from typing import Any + +_PUBLISHER: Any | None = None +_TOPIC_PATH = "" +_BRT = timezone(timedelta(hours=-3)) +_ISO_EVENT_TIMESTAMP_RE = re.compile( + r"^(\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2})(?:\.(\d{1,9}))?Z$" +) +_PUBSUB_EVENT_TIMESTAMP_RE = re.compile(r"^(\d{2}/\d{2}/\d{4} \d{2}:\d{2}:\d{2}),(\d{1,9})$") +_EXCLUDED_PAYLOAD_KEYS = {"original_message_id"} + + +def _config() -> tuple[str, str] | None: + project_id = (os.getenv("GCP_PROJECT_ID", "") or "").strip() + topic = (os.getenv("AGENT_PUBSUB_TOPIC", "") or "").strip() + if not project_id or not topic: + return None + return project_id, topic + + +def _topic_path(publisher: Any, project_id: str, topic: str) -> str: + if topic.startswith("projects/") and "/topics/" in topic: + return topic + return str(publisher.topic_path(project_id, topic)) + + +def _publisher() -> tuple[Any, str] | None: + global _PUBLISHER, _TOPIC_PATH + + config = _config() + if config is None: + return None + + if _PUBLISHER is not None and _TOPIC_PATH: + return _PUBLISHER, _TOPIC_PATH + + project_id, topic = config + try: + from google.cloud import pubsub_v1 + + _PUBLISHER = pubsub_v1.PublisherClient() + _TOPIC_PATH = _topic_path(_PUBLISHER, project_id, topic) + return _PUBLISHER, _TOPIC_PATH + except Exception as exc: + print(f"STRUCTURED_EVENT_PUBSUB_DISABLED | error={type(exc).__name__}: {exc}", file=sys.stderr) + return None + + +def _log_publish_error(future: Any) -> None: + try: + future.result() + except Exception as exc: + print(f"STRUCTURED_EVENT_PUBSUB_PUBLISH_FAIL | error={type(exc).__name__}: {exc}", file=sys.stderr) + + +def _format_pubsub_timestamp(value: Any) -> str: + raw = str(value or "").strip() + if not raw: + return "" + + pubsub_match = _PUBSUB_EVENT_TIMESTAMP_RE.match(raw) + if pubsub_match is not None: + return f"{pubsub_match.group(1)},{pubsub_match.group(2).ljust(9, '0')[:9]}" + + iso_match = _ISO_EVENT_TIMESTAMP_RE.match(raw) + if iso_match is None: + return raw + + dt = datetime.strptime(iso_match.group(1), "%Y-%m-%dT%H:%M:%S") + local_dt = dt.replace(tzinfo=timezone.utc).astimezone(_BRT) + fractional_ns = (iso_match.group(2) or "").ljust(9, "0")[:9] + return f"{local_dt:%d/%m/%Y %H:%M:%S},{fractional_ns}" + + +def _pubsub_payload(event: Mapping[str, Any]) -> dict[str, Any]: + payload = {key: value for key, value in event.items() if str(key) not in _EXCLUDED_PAYLOAD_KEYS} + + if "session_id" in payload: + if payload.get("sessionId") in (None, ""): + payload["sessionId"] = payload.get("session_id") + del payload["session_id"] + + if "dat_hora_fim" in payload: + if payload.get("dat_hora_termino") in (None, ""): + payload["dat_hora_termino"] = payload.get("dat_hora_fim") + del payload["dat_hora_fim"] + + if "dat_hora_inicio" in payload: + payload["dat_hora_inicio"] = _format_pubsub_timestamp(payload.get("dat_hora_inicio")) + if "dat_hora_termino" in payload: + payload["dat_hora_termino"] = _format_pubsub_timestamp(payload.get("dat_hora_termino")) + + return payload + + +def publish_structured_event(event: Mapping[str, Any]) -> None: + publisher_config = _publisher() + if publisher_config is None: + return + + publisher, topic_path = publisher_config + try: + payload = _pubsub_payload(event) + payload_json = json.dumps(payload, ensure_ascii=False).encode("utf-8") + future = publisher.publish( + topic_path, + payload_json, + tipo_evento=str(payload.get("tipo_evento") or ""), + dat_hora_inicio=str(payload.get("dat_hora_inicio") or ""), + dat_hora_termino=str(payload.get("dat_hora_termino") or ""), + sessionId=str(payload.get("sessionId") or ""), + callid=str(payload.get("callid") or ""), + nome_agente=str(payload.get("nome_agente") or ""), + cod_ani=str(payload.get("cod_ani") or ""), + http_cod_status=str(payload.get("http_cod_status") or ""), + http_cod_desc=str(payload.get("http_cod_desc") or ""), + ) + future.add_done_callback(_log_publish_error) + except Exception as exc: + print(f"STRUCTURED_EVENT_PUBSUB_PUBLISH_FAIL | error={type(exc).__name__}: {exc}", file=sys.stderr) diff --git a/src/app/utils/stt_audio_upload.py b/src/app/utils/stt_audio_upload.py new file mode 100644 index 0000000..7c66671 --- /dev/null +++ b/src/app/utils/stt_audio_upload.py @@ -0,0 +1,408 @@ +from __future__ import annotations + +import asyncio +import hashlib +import logging +import os +import re +import threading +import time +import weakref +from dataclasses import dataclass +from datetime import datetime +from typing import Any, Mapping + +from app.utils.logging import log_flow_event + +logger = logging.getLogger(__name__) + +_QUEUE_SIZE = 64 +_MAX_WAV_BYTES = 5 * 1024 * 1024 +_MAX_OBJECT_NAME_BYTES = 512 +_MAX_PART_CHARS = 64 +_SAFE_PART_RE = re.compile(r"[^A-Za-z0-9_.-]+") +_OCI_AUTH_MODE_LOCAL = "local" +_OCI_AUTH_MODE_OKE_WORKLOAD_IDENTITY = "oke_workload_identity" + + +@dataclass(frozen=True, slots=True) +class STTAudioUploadItem: + object_name: str + wav_bytes: bytes + req_id: str + message_id: str + session_id: str + metadata: Mapping[str, Any] + logger: Any | None = None + timeline: Any | None = None + + +@dataclass(frozen=True, slots=True) +class OCIUploadConfig: + auth_mode: str + region: str + bucket: str + namespace: str + user: str = "" + key_content: str = "" + fingerprint: str = "" + tenancy: str = "" + + +def _safe_path_part(value: Any, *, fallback: str) -> str: + raw = str(value or "").strip() + if not raw: + raw = fallback + safe = _SAFE_PART_RE.sub("_", raw).strip("._-") or fallback + if len(safe) <= _MAX_PART_CHARS: + return safe + + digest = hashlib.sha1(safe.encode("utf-8")).hexdigest()[:12] + prefix_len = max(1, _MAX_PART_CHARS - len(digest) - 1) + return f"{safe[:prefix_len].rstrip('._-')}-{digest}" + + +def _object_name_within_limit(object_name: str) -> str: + if len(object_name.encode("utf-8")) <= _MAX_OBJECT_NAME_BYTES: + return object_name + + digest = hashlib.sha1(object_name.encode("utf-8")).hexdigest() + today = datetime.now().strftime("%Y-%m-%d") + return f"{today}/oversize/{digest}.wav" + + +def build_stt_vad_object_name( + *, + session_id: Any, + message_id: Any, + req_id: Any, + now: datetime | None = None, +) -> str: + dt = now or datetime.now() + session_part = _safe_path_part(session_id, fallback="unknown-session") + message_part = _safe_path_part(message_id, fallback="unknown-message") + object_name = ( + f"{dt:%Y-%m-%d}/{session_part}/{message_part}.wav" + ) + return _object_name_within_limit(object_name) + + +def _metadata_value(source: Any, *keys: str) -> str: + if isinstance(source, Mapping): + lowered = {str(key).lower(): value for key, value in source.items()} + for key in keys: + value = source.get(key) + if value not in (None, ""): + return str(value).strip() + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return str(value).strip() + return "" + + for key in keys: + value = getattr(source, key, None) + if value not in (None, ""): + return str(value).strip() + return "" + + +def _oci_upload_config_from_env() -> OCIUploadConfig | None: + auth_mode = ( + _OCI_AUTH_MODE_LOCAL + if os.getenv("OCI_AUTH_MODE", "").strip().lower() == _OCI_AUTH_MODE_LOCAL + else _OCI_AUTH_MODE_OKE_WORKLOAD_IDENTITY + ) + shared_values = { + "region": os.getenv("BUCKET_REGION", "").strip(), + "bucket": os.getenv("BUCKET_NAME", "").strip(), + "namespace": os.getenv("BUCKET_NAMESPACE", "").strip(), + } + if not all(shared_values.values()): + return None + + if auth_mode != _OCI_AUTH_MODE_LOCAL: + return OCIUploadConfig(auth_mode=auth_mode, **shared_values) + + local_values = { + "user": os.getenv("BUCKET_USER_ID", "").strip(), + "key_content": os.getenv("BUCKET_PRIVATE_KEY", "").replace("\\n", "\n").strip(), + "fingerprint": os.getenv("BUCKET_FINGERPRINT", "").strip(), + "tenancy": os.getenv("BUCKET_TENANCY_ID", "").strip(), + } + if not all(local_values.values()): + return None + return OCIUploadConfig(auth_mode=auth_mode, **shared_values, **local_values) + + +_OCI_CLIENT: Any | None = None +_OCI_CLIENT_CONFIG: OCIUploadConfig | None = None +_OCI_CLIENT_LOCK = threading.Lock() +# ObjectStorageClient owns a requests Session and must not be used concurrently +# by the STT and entire-call upload workers. +_OCI_UPLOAD_LOCK = threading.Lock() + + +def _oci_client(config: OCIUploadConfig) -> Any: + global _OCI_CLIENT, _OCI_CLIENT_CONFIG + + with _OCI_CLIENT_LOCK: + if _OCI_CLIENT is not None and _OCI_CLIENT_CONFIG == config: + return _OCI_CLIENT + + from oci.object_storage import ObjectStorageClient + + if config.auth_mode == _OCI_AUTH_MODE_LOCAL: + client_config = { + "user": config.user, + "key_content": config.key_content, + "fingerprint": config.fingerprint, + "tenancy": config.tenancy, + "region": config.region, + } + _OCI_CLIENT = ObjectStorageClient(config=client_config) + else: + from oci.auth.signers import get_oke_workload_identity_resource_principal_signer + + signer = get_oke_workload_identity_resource_principal_signer() + _OCI_CLIENT = ObjectStorageClient( + config={"region": config.region}, + signer=signer, + ) + _OCI_CLIENT_CONFIG = config + session = getattr(getattr(_OCI_CLIENT, "base_client", None), "session", None) + if session is not None: + session.trust_env = False + return _OCI_CLIENT + + +def _reset_oci_client() -> None: + """Discard a client whose HTTP connection pool may be in an invalid state.""" + global _OCI_CLIENT, _OCI_CLIENT_CONFIG + + with _OCI_CLIENT_LOCK: + client = _OCI_CLIENT + _OCI_CLIENT = None + _OCI_CLIENT_CONFIG = None + + session = getattr(getattr(client, "base_client", None), "session", None) + close = getattr(session, "close", None) + if callable(close): + try: + close() + except Exception: + logger.debug("failed closing discarded OCI client session", exc_info=True) + + +def _put_object_sync(item: STTAudioUploadItem, config: OCIUploadConfig) -> Any: + with _OCI_UPLOAD_LOCK: + client = _oci_client(config) + return client.put_object( + namespace_name=config.namespace, + bucket_name=config.bucket, + object_name=item.object_name, + put_object_body=item.wav_bytes, + ) + + +class _STTAudioUploadWorker: + def __init__(self, loop: asyncio.AbstractEventLoop) -> None: + self._loop = loop + self._queue: asyncio.Queue[STTAudioUploadItem] = asyncio.Queue(maxsize=_QUEUE_SIZE) + self._task: asyncio.Task[None] | None = None + + def enqueue(self, item: STTAudioUploadItem) -> bool: + if self._task is None or self._task.done(): + self._task = self._loop.create_task( + self._run(), + name="stt-audio-upload-worker", + ) + try: + self._queue.put_nowait(item) + return True + except asyncio.QueueFull: + return False + + async def _run(self) -> None: + while True: + item = await self._queue.get() + try: + await _upload_item(item) + except Exception: + log_flow_event( + item.logger or logger, + "stt_audio_upload_failed", + request_id=item.req_id, + message_id=item.message_id, + session_id=item.session_id, + object_name=item.object_name, + bytes=len(item.wav_bytes), + **_log_metadata_fields(item), + ) + _emit_timeline( + item, + "stt_audio_upload_failed", + message_id=item.message_id, + object_name=item.object_name, + bytes=len(item.wav_bytes), + **_log_metadata_fields(item), + ) + (item.logger or logger).exception( + "[stt][audio_upload_failed] req_id=%s object_name=%s", + item.req_id, + item.object_name, + ) + finally: + self._queue.task_done() + + +_WORKERS: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, _STTAudioUploadWorker] = ( + weakref.WeakKeyDictionary() +) +_WORKERS_LOCK = threading.Lock() + + +def _worker_for_loop(loop: asyncio.AbstractEventLoop) -> _STTAudioUploadWorker: + with _WORKERS_LOCK: + worker = _WORKERS.get(loop) + if worker is None: + worker = _STTAudioUploadWorker(loop) + _WORKERS[loop] = worker + return worker + + +def _emit_timeline(item: STTAudioUploadItem, event: str, **fields: Any) -> None: + if item.timeline is None: + return + try: + item.timeline.emit(event, request_id=item.req_id, **fields) + except Exception: + pass + + +def _log_metadata_fields(item: STTAudioUploadItem) -> dict[str, Any]: + return { + str(key): value + for key, value in item.metadata.items() + if value not in (None, "") + } + + +async def _upload_item(item: STTAudioUploadItem) -> None: + started = time.perf_counter() + config = _oci_upload_config_from_env() + if config is None: + raise RuntimeError("OCI Object Storage upload config is missing") + + response = await asyncio.to_thread(_put_object_sync, item, config) + + duration_ms = round((time.perf_counter() - started) * 1000) + log_flow_event( + item.logger or logger, + "stt_audio_upload_done", + request_id=item.req_id, + message_id=item.message_id, + session_id=item.session_id, + object_name=item.object_name, + bytes=len(item.wav_bytes), + duration_ms=duration_ms, + bucket=config.bucket, + namespace=config.namespace, + oci_status=getattr(response, "status", None), + **_log_metadata_fields(item), + ) + _emit_timeline( + item, + "stt_audio_upload_done", + message_id=item.message_id, + object_name=item.object_name, + bytes=len(item.wav_bytes), + duration_ms=duration_ms, + bucket=config.bucket, + namespace=config.namespace, + oci_status=getattr(response, "status", None), + **_log_metadata_fields(item), + ) + + +def enqueue_stt_vad_audio_upload( + *, + wav_bytes: bytes, + req_id: str, + message_id: str, + structured_log_context: Any = None, + metadata: Mapping[str, Any] | None = None, + logger_override: Any | None = None, + timeline: Any | None = None, +) -> str | None: + if _oci_upload_config_from_env() is None: + return None + + active_logger = logger_override or logger + session_id = _metadata_value(structured_log_context, "session_id", "sessionId") + object_name = build_stt_vad_object_name( + session_id=session_id, + message_id=message_id, + req_id=req_id, + ) + + if len(wav_bytes) > _MAX_WAV_BYTES: + log_flow_event( + active_logger, + "stt_audio_upload_drop", + request_id=req_id, + message_id=message_id, + session_id=session_id, + object_name=object_name, + reason="too_large", + bytes=len(wav_bytes), + ) + return object_name + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + log_flow_event( + active_logger, + "stt_audio_upload_drop", + request_id=req_id, + message_id=message_id, + session_id=session_id, + object_name=object_name, + reason="no_running_loop", + ) + return object_name + + item = STTAudioUploadItem( + object_name=object_name, + wav_bytes=wav_bytes, + req_id=req_id, + message_id=message_id, + session_id=session_id, + metadata=metadata or {}, + logger=active_logger, + timeline=timeline, + ) + queued = _worker_for_loop(loop).enqueue(item) + if not queued: + log_flow_event( + active_logger, + "stt_audio_upload_drop", + request_id=req_id, + message_id=message_id, + session_id=session_id, + object_name=object_name, + reason="queue_full", + bytes=len(wav_bytes), + ) + return object_name + + log_flow_event( + active_logger, + "stt_audio_upload_enqueued", + request_id=req_id, + message_id=message_id, + session_id=session_id, + object_name=object_name, + bytes=len(wav_bytes), + ) + return object_name diff --git a/src/app/utils/turn_ids.py b/src/app/utils/turn_ids.py new file mode 100644 index 0000000..583aec7 --- /dev/null +++ b/src/app/utils/turn_ids.py @@ -0,0 +1,162 @@ +from __future__ import annotations + +import threading +import uuid +from collections import defaultdict, deque +from dataclasses import dataclass +from typing import Any, Mapping + + +_LOCK = threading.Lock() +_PENDING_TRANSCRIPTS: dict[str, deque["RegisteredTurn"]] = defaultdict(deque) +_STARTED_TURN_MESSAGE_IDS: dict[str, str] = {} +_MAX_PENDING_TRANSCRIPTS = 32 + + +@dataclass(frozen=True, slots=True) +class RegisteredTurn: + message_id: str + transcription: str + text: str + + +def _string_or_empty(value: Any) -> str: + return str(value or "").strip() + + +def _mapping_value(payload: Mapping[str, Any], *keys: str) -> str: + for key in keys: + value = payload.get(key) + if value not in (None, ""): + return _string_or_empty(value) + + lowered = {str(key).lower(): value for key, value in payload.items()} + for key in keys: + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return _string_or_empty(value) + + return "" + + +def _source_value(source: Any, *keys: str) -> str: + if isinstance(source, Mapping): + return _mapping_value(source, *keys) + + for key in keys: + value = getattr(source, key, None) + if value not in (None, ""): + return _string_or_empty(value) + return "" + + +def _router_callid(source: Any) -> str: + day = _source_value(source, "RouterCallKeyDay", "routerCallKeyDay", "router_call_key_day") + key = _source_value(source, "RouterCallKey", "routerCallKey", "router_call_key") + if day or key: + return f"{day}{key}" + return "" + + +def turn_message_key(source: Any) -> str: + return ( + _source_value(source, "session_id", "sessionId") + or _router_callid(source) + or _source_value(source, "callid", "call_id") + or _source_value(source, "callIdGed", "call_id_ged", "message_id_seed") + or _source_value(source, "protocol_id", "protocolId", "protocolo", "protocolNumber", "protocol") + or "turn" + ) + + +def reset_turn_message_sequence(source: Any, *, clear_pending: bool = False) -> None: + if not clear_pending: + return + + key = turn_message_key(source) + with _LOCK: + _PENDING_TRANSCRIPTS.pop(key, None) + _STARTED_TURN_MESSAGE_IDS.pop(key, None) + + +def next_turn_message_id(source: Any) -> str: + del source + return str(uuid.uuid4()) + + +def register_started_turn_message_id(source: Any, message_id: str) -> None: + normalized_message_id = _string_or_empty(message_id) + if not normalized_message_id: + return + + key = turn_message_key(source) + with _LOCK: + _STARTED_TURN_MESSAGE_IDS[key] = normalized_message_id + + +def peek_started_turn_message_id(source: Any) -> str: + key = turn_message_key(source) + with _LOCK: + return _STARTED_TURN_MESSAGE_IDS.get(key, "") + + +def clear_started_turn_message_id(source: Any, message_id: str = "") -> None: + key = turn_message_key(source) + normalized_message_id = _string_or_empty(message_id) + with _LOCK: + if not normalized_message_id: + _STARTED_TURN_MESSAGE_IDS.pop(key, None) + return + + if _STARTED_TURN_MESSAGE_IDS.get(key) == normalized_message_id: + _STARTED_TURN_MESSAGE_IDS.pop(key, None) + + +def register_transcribed_turn( + source: Any, + *, + message_id: str, + transcription: str = "", + text: str = "", +) -> None: + normalized_message_id = _string_or_empty(message_id) + if not normalized_message_id: + return + + key = turn_message_key(source) + turn = RegisteredTurn( + message_id=normalized_message_id, + transcription=_string_or_empty(transcription), + text=_string_or_empty(text), + ) + with _LOCK: + queue = _PENDING_TRANSCRIPTS[key] + queue.append(turn) + while len(queue) > _MAX_PENDING_TRANSCRIPTS: + queue.popleft() + + +def consume_transcribed_turn( + source: Any, + *, + transcription: str = "", + text: str = "", +) -> RegisteredTurn | None: + key = turn_message_key(source) + normalized_transcription = _string_or_empty(transcription) + normalized_text = _string_or_empty(text) + + with _LOCK: + queue = _PENDING_TRANSCRIPTS.get(key) + if not queue: + return None + + for turn in tuple(queue): + if normalized_transcription and turn.transcription == normalized_transcription: + queue.remove(turn) + return turn + if normalized_text and turn.text == normalized_text: + queue.remove(turn) + return turn + + return queue.popleft() diff --git a/src/app/ws_gateway/call_config.py b/src/app/ws_gateway/call_config.py new file mode 100644 index 0000000..585ee8a --- /dev/null +++ b/src/app/ws_gateway/call_config.py @@ -0,0 +1,9 @@ +from __future__ import annotations + +from typing import Any, Dict, Mapping + +from app.common.call_config import normalize_call_config + + +def build_call_config(payload: Mapping[str, Any] | None) -> Dict[str, Any]: + return normalize_call_config(payload) diff --git a/src/app/ws_gateway/fake_remote_agent.py b/src/app/ws_gateway/fake_remote_agent.py new file mode 100644 index 0000000..0387a63 --- /dev/null +++ b/src/app/ws_gateway/fake_remote_agent.py @@ -0,0 +1,149 @@ +from __future__ import annotations + +from typing import Any, Mapping + + +_CONTA_OPENING = ( + "Olá! Eu sou a Especialista em Contas e vou ajudar você a entender a sua " + "fatura. Posso explicar valores, detalhar serviços e itens eventuais, " + "identificar cobranças que você não reconhece e, se for o caso, realizar " + "ajustes necessários ou solicitações relacionadas à sua conta. Então vamos " + "lá, me conte o que você gostaria de entender ou resolver na sua conta." +) + + +def _as_str(value: Any) -> str: + return str(value or "").strip() + + +def _normalize_agent_name(value: Any) -> str: + raw = _as_str(value).lower() + aliases = { + "conta": "conta", + "contas": "conta", + "ofert": "oferta", + "oferta": "oferta", + "ofertas": "oferta", + "cobra": "cobranca", + "cobranca": "cobranca", + "cobranca": "cobranca", + } + return aliases.get(raw, raw or "conta") + + +def _payload_body(payload: Mapping[str, Any]) -> Mapping[str, Any]: + body = payload.get("payload") + if isinstance(body, Mapping): + return body + return payload + + +def _extract_agent_name(payload: Mapping[str, Any]) -> str: + body = _payload_body(payload) + return _normalize_agent_name( + body.get("agent") + or payload.get("agent") + or payload.get("_agent") + ) + + +def _extract_stage(payload: Mapping[str, Any]) -> str: + body = _payload_body(payload) + return _as_str( + body.get("stage") + or payload.get("stage") + or payload.get("_stage") + ).upper() or "PRESENTATION" + + +def _extract_user_text(payload: Mapping[str, Any]) -> str: + body = _payload_body(payload) + for key in ("message", "text", "utterance", "transcript", "content"): + text = _as_str(body.get(key) or payload.get(key)) + if text: + return text + return "" + + +def _is_end_request(payload: Mapping[str, Any]) -> bool: + body = _payload_body(payload) + if _as_str(payload.get("action")).lower() == "end": + return True + if _as_str(payload.get("type")).lower() == "end": + return True + if _as_str(body.get("type")).lower() == "end": + return True + return False + + +def _next_stage(current_stage: str, text: str, *, is_end: bool) -> str: + lowered = _as_str(text).lower() + if is_end or any(token in lowered for token in ("encerrar", "finalizar", "obrigado", "tchau")): + return "DONE" + + stage = (current_stage or "PRESENTATION").upper() + if stage in {"", "INTRO", "PRESENTATION"}: + return "ARGUMENTATION" + if stage == "ARGUMENTATION": + return "FORMALIZATION" + if stage == "FORMALIZATION": + return "DONE" + return stage + + +def _reply_text(agent_name: str, next_stage: str, user_text: str) -> str: + transcript = _as_str(user_text) + + if next_stage == "DONE": + if agent_name == "conta": + return "Atendimento simulado de conta encerrado. Obrigado." + if agent_name == "oferta": + return "Atendimento simulado de oferta encerrado. Obrigado." + return "Atendimento simulado de cobrança encerrado. Obrigado." + + if agent_name == "conta": + if not transcript: + return _CONTA_OPENING + base = "Simulação de conta ativa. Posso seguir com os detalhes da fatura." + elif agent_name == "oferta": + base = "Simulação de oferta ativa. Tenho uma oferta pronta para você." + else: + base = "Simulação de cobrança ativa. Posso continuar com a negociação." + + if transcript: + return f"{base} Texto transcrito no STT: {transcript}." + return base + + +def build_fake_remote_agent_response(payload: Mapping[str, Any]) -> dict[str, Any]: + agent_name = _extract_agent_name(payload) + current_stage = _extract_stage(payload) + user_text = _extract_user_text(payload) + is_end = _is_end_request(payload) + next_stage = _next_stage(current_stage, user_text, is_end=is_end) + content = _reply_text(agent_name, next_stage, user_text) + + if agent_name == "conta": + action = _as_str(payload.get("action")).lower() or "chat" + result_payload = { + "type": "final", + "content": content, + "tool_calls": [], + } + if next_stage == "DONE": + result_payload["result"] = [{"status": "ok", "reason": "fake_done"}] + return { + "type": "result", + "action": action or "chat", + "stage": next_stage, + "result": result_payload, + } + + response_type = "done" if next_stage == "DONE" else "final" + result = [{"status": "ok", "reason": "fake_done"}] if next_stage == "DONE" else [] + return { + "type": response_type, + "stage": next_stage, + "text": content, + "result": result, + } diff --git a/src/app/ws_gateway/main.py b/src/app/ws_gateway/main.py new file mode 100644 index 0000000..529aaa6 --- /dev/null +++ b/src/app/ws_gateway/main.py @@ -0,0 +1,2100 @@ +# ws_bridge.py +import asyncio +import json +import os +import time +import uuid +import audioop +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable, Dict, Mapping, Optional + +from dotenv import load_dotenv +from fastapi import FastAPI, WebSocket, WebSocketDisconnect +from fastapi.responses import FileResponse, JSONResponse + +from livekit import rtc +from livekit.api import AccessToken, VideoGrants +from livekit import api as lkapi + +from app.common.call_config import ( + resolve_agent_backend_name, + resolve_stt_overrides, + resolve_tts_overrides, + resolve_ws_overrides, +) +from app.ws_gateway.fake_remote_agent import build_fake_remote_agent_response +from app.ws_gateway.session_audio import ( + enable_client_audio, + notify_client_audio_enabled, + stream_agent_audio_when_ready, + watch_mock_stop_after_first_audio, +) +from app.ws_gateway.session_bootstrap import build_bridge_session_bootstrap +from app.ws_gateway.session_lifecycle import ( + RoomLifecycleState, + close_websocket, + register_room_lifecycle_handlers, + watch_call_done, +) +from app.ws_gateway.readiness import ( + BridgeReadinessManager, + build_bridge_failed_stop_message, + build_capacity_stop_message, +) +from app.ws_gateway.session_start import recv_start_message +from app.utils.call_timeline import CallTimeline +from app.utils.audio_backlog import ShedResult, rms_threshold_from_dbfs, shed_queue_backlog +from app.utils.background import AudioActivity, BridgeOutputStats, ws_out_loop +from app.utils.full_call_recording import create_entire_call_recorder_from_env +from app.utils.pcm import pcm16_dbfs +from app.utils.logging import log_flow_event, setup_minimal_logging + +logger = setup_minimal_logging() +app = FastAPI() +READINESS_GATE = BridgeReadinessManager() + +SAMPLE_RATE = 16000 +CHANNELS = 1 +FRAME_MS = 20 +SAMPLES_PER_FRAME = int(SAMPLE_RATE * FRAME_MS / 1000) # 320 +BYTES_PER_FRAME = SAMPLES_PER_FRAME * 2 * CHANNELS # 640 + +HOLD_SILENCE_DBFS = float(os.getenv("HOLD_SILENCE_DBFS", "-50")) +RMS_TH = int(32768 * (10 ** (HOLD_SILENCE_DBFS / 20.0))) +SILENCE_CHECK_EVERY = int(os.getenv("SILENCE_CHECK_EVERY", "5")) # checa a cada 5 frames (100ms) + +# --- ENV / AMBIENTE --- +APP_ENV = (os.getenv("APP_ENV") or os.getenv("ENV") or "dev").strip().lower() +DOTENV_PATH = os.getenv("DOTENV_PATH", f".env.{APP_ENV}") + +AGENT_BASE_NAME = os.getenv("AGENT_BASE_NAME", "ws-voice-agent").strip() +AGENT_NAME = os.getenv("AGENT_NAME", f"{AGENT_BASE_NAME}-{APP_ENV}").strip() + +LK_CATCHUP_ENABLED = os.getenv("LK_CATCHUP_ENABLED", "1") == "1" +LK_CATCHUP_KEEP_MS = int(os.getenv("LK_CATCHUP_KEEP_MS", "1000")) +LK_CATCHUP_MAX_FAST_FRAMES = int(os.getenv("LK_CATCHUP_MAX_FAST_FRAMES", "4000")) + +# --- Marcador de "fala ativa" do ÁUDIO DO AGENTE (só telemetria/idle; NÃO filtra áudio) --- +AGENT_OUT_SPEECH_DBFS = float(os.getenv("AGENT_OUT_SPEECH_DBFS", "-52")) # acima disso = fala ativa +AGENT_OUT_SPEECH_RMS_TH = int(32768 * (10 ** (AGENT_OUT_SPEECH_DBFS / 20.0))) +AGENT_OUT_RMS_EVERY = int(os.getenv("AGENT_OUT_RMS_EVERY", "1")) # 1 = checa todo frame +FLOW_AUDIO_BURST_GAP_S = float(os.getenv("FLOW_AUDIO_BURST_GAP_S", "0.8")) +FLOW_LOG_AUDIO_BURST_END = ( + os.getenv("FLOW_LOG_AUDIO_BURST_END", "0").strip().lower() + in {"1", "true", "yes", "on"} +) +VOICE_CLIENT_HTML = Path(__file__).with_name("voice_client.html") + + +def _load_runtime_env() -> None: + if DOTENV_PATH: + load_dotenv(DOTENV_PATH, override=False) + load_dotenv(override=False) + + +def _env_bool(name: str, default: bool) -> bool: + raw = os.getenv(name) + if raw is None: + return default + return raw.strip().lower() in {"1", "true", "yes", "on"} + + +def _env_int(name: str, default: int) -> int: + try: + return int(os.getenv(name, str(default))) + except ValueError: + return default + + +def _env_float(name: str, default: float) -> float: + try: + return float(os.getenv(name, str(default))) + except ValueError: + return default + + +def _frames_from_ms(value_ms: int, *, minimum: int = 1) -> int: + value_ms = max(0, int(value_ms)) + return max(minimum, (value_ms + FRAME_MS - 1) // FRAME_MS) + + +@dataclass(frozen=True) +class AudioInputBacklogConfig: + shed_enabled: bool + shed_threshold_ms: int + shed_keep_ms: int + latency_metrics_enabled: bool + latency_alert_ms: int + latency_log_interval_s: float + livekit_source_queue_size_ms: int + livekit_source_clear_on_shed: bool + config_source: str = "env" + call_config_overrides: tuple[str, ...] = () + # Descarte guiado por energia: no regime de drift (excesso pequeno) so + # descarta silencio; so cai no descarte cego quando o excesso passa de + # energy_shed_max_excess_ms (regime de rajada). + energy_shed_enabled: bool = True + energy_shed_max_excess_ms: int = 1000 + silence_dbfs: float = -50.0 + + @property + def shed_threshold_frames(self) -> int: + return _frames_from_ms(self.shed_threshold_ms) + + @property + def shed_keep_frames(self) -> int: + return min( + self.shed_threshold_frames, + _frames_from_ms(self.shed_keep_ms), + ) + + @property + def silence_rms_threshold(self) -> int: + return rms_threshold_from_dbfs(self.silence_dbfs) + + +def _call_config_has_value(overrides: Mapping[str, Any], key: str) -> bool: + return overrides.get(key) not in (None, "") + + +def _config_bool( + overrides: Mapping[str, Any], + key: str, + env_name: str, + default: bool, +) -> bool: + env_value = _env_bool(env_name, default) + if not _call_config_has_value(overrides, key): + return env_value + return str(overrides.get(key)).strip().lower() in {"1", "true", "yes", "on"} + + +def _config_int( + overrides: Mapping[str, Any], + key: str, + env_name: str, + default: int, +) -> int: + env_value = _env_int(env_name, default) + if not _call_config_has_value(overrides, key): + return env_value + try: + return int(str(overrides.get(key)).strip()) + except (TypeError, ValueError): + return env_value + + +def _config_float( + overrides: Mapping[str, Any], + key: str, + env_name: str, + default: float, +) -> float: + env_value = _env_float(env_name, default) + if not _call_config_has_value(overrides, key): + return env_value + try: + return float(str(overrides.get(key)).strip()) + except (TypeError, ValueError): + return env_value + + +def audio_input_backlog_config_from_env( + call_config: Optional[Mapping[str, Any]] = None, +) -> AudioInputBacklogConfig: + ws_overrides = resolve_ws_overrides(call_config) + override_fields = tuple( + key + for key in ( + "audio_in_backlog_shed_enabled", + "audio_in_backlog_shed_threshold_ms", + "audio_in_backlog_shed_keep_ms", + "audio_in_latency_metrics_enabled", + "audio_in_latency_alert_ms", + "audio_in_latency_log_interval_s", + "livekit_audio_source_queue_size_ms", + "livekit_audio_source_clear_on_shed", + "audio_in_backlog_energy_shed_enabled", + "audio_in_backlog_energy_shed_max_excess_ms", + "audio_in_backlog_silence_dbfs", + ) + if _call_config_has_value(ws_overrides, key) + ) + shed_enabled = _config_bool( + ws_overrides, + "audio_in_backlog_shed_enabled", + "AUDIO_IN_BACKLOG_SHED_ENABLED", + True, + ) + default_source_queue_ms = 500 if shed_enabled else 5000 + return AudioInputBacklogConfig( + shed_enabled=shed_enabled, + shed_threshold_ms=max( + FRAME_MS, + _config_int( + ws_overrides, + "audio_in_backlog_shed_threshold_ms", + "AUDIO_IN_BACKLOG_SHED_THRESHOLD_MS", + 500, + ), + ), + shed_keep_ms=max( + FRAME_MS, + _config_int( + ws_overrides, + "audio_in_backlog_shed_keep_ms", + "AUDIO_IN_BACKLOG_SHED_KEEP_MS", + 300, + ), + ), + latency_metrics_enabled=_config_bool( + ws_overrides, + "audio_in_latency_metrics_enabled", + "AUDIO_IN_LATENCY_METRICS_ENABLED", + True, + ), + latency_alert_ms=max( + 0, + _config_int( + ws_overrides, + "audio_in_latency_alert_ms", + "AUDIO_IN_LATENCY_ALERT_MS", + 1000, + ), + ), + latency_log_interval_s=max( + 0.1, + _config_float( + ws_overrides, + "audio_in_latency_log_interval_s", + "AUDIO_IN_LATENCY_LOG_INTERVAL_S", + 15.0, + ), + ), + livekit_source_queue_size_ms=max( + FRAME_MS, + _config_int( + ws_overrides, + "livekit_audio_source_queue_size_ms", + "LIVEKIT_AUDIO_SOURCE_QUEUE_SIZE_MS", + default_source_queue_ms, + ), + ), + livekit_source_clear_on_shed=_config_bool( + ws_overrides, + "livekit_audio_source_clear_on_shed", + "LIVEKIT_AUDIO_SOURCE_CLEAR_ON_SHED", + shed_enabled, + ), + energy_shed_enabled=_config_bool( + ws_overrides, + "audio_in_backlog_energy_shed_enabled", + "AUDIO_IN_BACKLOG_ENERGY_SHED_ENABLED", + True, + ), + energy_shed_max_excess_ms=max( + 0, + _config_int( + ws_overrides, + "audio_in_backlog_energy_shed_max_excess_ms", + "AUDIO_IN_BACKLOG_ENERGY_SHED_MAX_EXCESS_MS", + 1000, + ), + ), + silence_dbfs=_config_float( + ws_overrides, + "audio_in_backlog_silence_dbfs", + "AUDIO_IN_BACKLOG_SILENCE_DBFS", + -60.0, + ), + config_source="call_config" if override_fields else "env", + call_config_overrides=override_fields, + ) + + +class AudioInputLatencyTracker: + def __init__( + self, + config: AudioInputBacklogConfig, + *, + timeline: Optional[CallTimeline] = None, + call_context: Optional[Mapping[str, Any]] = None, + debug_event_publisher: Optional[Callable[[str, Mapping[str, Any]], None]] = None, + ) -> None: + self.config = config + self.timeline = timeline + self.call_context = dict(call_context or {}) + self.debug_event_publisher = debug_event_publisher + self.enabled_at: float | None = None + self.last_latency_log_at = 0.0 + self.total_dropped_frames = 0 + self.peak_excess_ms = 0 + self.peak_raw_excess_ms = 0 + + def mark_enabled(self, now: Optional[float] = None) -> None: + if self.enabled_at is None: + self.enabled_at = time.monotonic() if now is None else now + self.last_latency_log_at = 0.0 + + def snapshot( + self, + *, + frame_count: int, + now: Optional[float] = None, + queue_size_frames: int = 0, + ) -> Dict[str, int]: + if self.enabled_at is None: + self.mark_enabled(now) + now = time.monotonic() if now is None else now + enabled_at = self.enabled_at if self.enabled_at is not None else now + elapsed_ms = max(0, round((now - enabled_at) * 1000)) + frames_since_enable = max(0, int(frame_count)) + received_audio_ms = frames_since_enable * FRAME_MS + effective_frames = max(0, frames_since_enable - self.total_dropped_frames) + effective_audio_ms = effective_frames * FRAME_MS + raw_excess_ms = received_audio_ms - elapsed_ms + excess_ms = effective_audio_ms - elapsed_ms + raw_lag_ms = max(0, raw_excess_ms) + lag_ms = max(0, excess_ms) + self.peak_excess_ms = max(self.peak_excess_ms, excess_ms) + self.peak_raw_excess_ms = max(self.peak_raw_excess_ms, raw_excess_ms) + return { + "frames_since_enable": frames_since_enable, + "received_audio_ms": received_audio_ms, + "effective_frames_since_enable": effective_frames, + "effective_audio_ms": effective_audio_ms, + "elapsed_since_enable_ms": elapsed_ms, + "excess_ms": excess_ms, + "lag_ms": lag_ms, + "peak_excess_ms": self.peak_excess_ms, + "raw_excess_ms": raw_excess_ms, + "raw_lag_ms": raw_lag_ms, + "peak_raw_excess_ms": self.peak_raw_excess_ms, + "queue_frames": max(0, int(queue_size_frames)), + "queue_ms": max(0, int(queue_size_frames)) * FRAME_MS, + "total_dropped_frames": self.total_dropped_frames, + "total_dropped_ms": self.total_dropped_frames * FRAME_MS, + "saved_latency_ms": self.total_dropped_frames * FRAME_MS, + } + + def maybe_log_latency( + self, + *, + frame_count: int, + now: Optional[float] = None, + queue_size_frames: int = 0, + reason: str, + force: bool = False, + ) -> None: + if not self.config.latency_metrics_enabled: + return + now = time.monotonic() if now is None else now + fields = self.snapshot( + frame_count=frame_count, + now=now, + queue_size_frames=queue_size_frames, + ) + should_log = ( + force + or fields["lag_ms"] >= self.config.latency_alert_ms + or fields["queue_ms"] >= self.config.latency_alert_ms + or (now - self.last_latency_log_at) >= self.config.latency_log_interval_s + ) + if not should_log: + return + self.last_latency_log_at = now + payload = { + **self.call_context, + **fields, + "reason": reason, + "alert_ms": self.config.latency_alert_ms, + } + log_flow_event(logger, "audio_in_latency", **payload) + if self.timeline is not None: + self.timeline.emit("audio_in_latency", **payload) + + def record_shed( + self, + *, + source: str, + dropped_frames: int, + queue_before_frames: int, + queue_after_frames: int, + frame_count: int, + mode: str = "blind", + dropped_silent: int = 0, + dropped_voiced: int = 0, + now: Optional[float] = None, + ) -> None: + if dropped_frames <= 0: + return + now = time.monotonic() if now is None else now + self.total_dropped_frames += int(dropped_frames) + fields = self.snapshot( + frame_count=frame_count, + now=now, + queue_size_frames=queue_after_frames, + ) + payload = { + **self.call_context, + **fields, + "source": source, + "mode": mode, + "content_policy": "silence_only" if mode == "energy" else "oldest_frames", + "sync_action": "compress_silence" if mode == "energy" else "drop_audio_to_catch_up", + "dropped_frames": int(dropped_frames), + "dropped_ms": int(dropped_frames) * FRAME_MS, + "dropped_silent_frames": int(dropped_silent), + "dropped_silent_ms": int(dropped_silent) * FRAME_MS, + "dropped_voiced_frames": int(dropped_voiced), + "dropped_voiced_ms": int(dropped_voiced) * FRAME_MS, + "saved_ms": int(dropped_frames) * FRAME_MS, + "queue_before_frames": int(queue_before_frames), + "queue_before_ms": int(queue_before_frames) * FRAME_MS, + "queue_after_frames": int(queue_after_frames), + "queue_after_ms": int(queue_after_frames) * FRAME_MS, + "shed_threshold_ms": self.config.shed_threshold_ms, + "shed_keep_ms": self.config.shed_keep_ms, + } + log_flow_event(logger, "audio_in_latency_shed", **payload) + if self.timeline is not None: + self.timeline.emit("audio_in_latency_shed", **payload) + if self.debug_event_publisher is not None: + self.debug_event_publisher("bridge.audio_in.shed", payload) + + def record_livekit_source_clear( + self, + *, + source: str, + queued_ms: int, + frame_count: int, + queue_size_frames: int, + now: Optional[float] = None, + ) -> None: + now = time.monotonic() if now is None else now + fields = self.snapshot( + frame_count=frame_count, + now=now, + queue_size_frames=queue_size_frames, + ) + payload = { + **self.call_context, + **fields, + "source": source, + "livekit_queued_ms": max(0, int(queued_ms)), + } + log_flow_event(logger, "audio_in_livekit_source_cleared", **payload) + if self.timeline is not None: + self.timeline.emit("audio_in_livekit_source_cleared", **payload) + + +def _shed_audio_queue_backlog( + audio_q: asyncio.Queue[bytes], + *, + config: AudioInputBacklogConfig, + tracker: Optional[AudioInputLatencyTracker], + source: str, + frame_count: int, + now: Optional[float] = None, +) -> ShedResult: + if not config.shed_enabled: + return ShedResult() + queue_before = audio_q.qsize() + if queue_before <= config.shed_threshold_frames: + return ShedResult() + + keep_frames = config.shed_keep_frames + over_keep_ms = max(0, queue_before - keep_frames) * FRAME_MS + # Regime de drift (excesso pequeno): so descarta silencio, preservando fala. + # Regime de rajada (excesso grande): descarte cego para recuperar latencia. + blind = ( + not config.energy_shed_enabled + or over_keep_ms > config.energy_shed_max_excess_ms + ) + result = shed_queue_backlog( + audio_q, + keep_frames=keep_frames, + rms_threshold=config.silence_rms_threshold, + blind=blind, + ) + + if result.dropped and tracker is not None: + tracker.record_shed( + source=source, + dropped_frames=result.dropped, + dropped_silent=result.dropped_silent, + dropped_voiced=result.dropped_voiced, + mode=result.mode, + queue_before_frames=queue_before, + queue_after_frames=audio_q.qsize(), + frame_count=frame_count, + now=now, + ) + return result + + +def _audio_input_latency_burst_fields( + tracker: Optional[AudioInputLatencyTracker], + *, + frame_count: int, + now: float, + queue_size_frames: int, +) -> Dict[str, Any]: + if tracker is None or not tracker.config.latency_metrics_enabled: + return {} + fields = tracker.snapshot( + frame_count=frame_count, + now=now, + queue_size_frames=queue_size_frames, + ) + return { + "input_excess_ms": fields["excess_ms"], + "input_lag_ms": fields["lag_ms"], + "input_raw_excess_ms": fields["raw_excess_ms"], + "input_raw_lag_ms": fields["raw_lag_ms"], + "input_queue_ms": fields["queue_ms"], + "input_total_dropped_ms": fields["total_dropped_ms"], + } + + +def _maybe_clear_livekit_source_queue( + source: rtc.AudioSource, + *, + tracker: Optional[AudioInputLatencyTracker], + reason: str, + frame_count: int, + queue_size_frames: int, + now: Optional[float] = None, +) -> int: + if tracker is None or not tracker.config.livekit_source_clear_on_shed: + return 0 + clear_queue = getattr(source, "clear_queue", None) + if not callable(clear_queue): + return 0 + try: + queued_ms = max(0, round(float(getattr(source, "queued_duration", 0.0)) * 1000)) + except Exception: + queued_ms = 0 + if queued_ms <= 0: + return 0 + try: + clear_queue() + except Exception: + logger.debug("LIVEKIT_AUDIO_SOURCE_CLEAR_FAIL", exc_info=True) + return 0 + tracker.record_livekit_source_clear( + source=reason, + queued_ms=queued_ms, + frame_count=frame_count, + queue_size_frames=queue_size_frames, + now=now, + ) + return queued_ms + + +def _is_new_audio_burst(now: float, last_audio_at: float) -> bool: + return last_audio_at <= 0.0 or (now - last_audio_at) >= FLOW_AUDIO_BURST_GAP_S + + +class FlowAudioBurstLogger: + def __init__(self, *, start_step: str, end_step: str) -> None: + self.start_step = start_step + self.end_step = end_step + self.seq = 0 + self.active = False + self.first_frame = 0 + self.last_frame = 0 + self.last_audio_at = 0.0 + self.active_frames = 0 + self.peak_rms = 0 + self.peak_dbfs = -120.0 + + def maybe_end(self, *, now: float, reason: str) -> None: + if self.active and _is_new_audio_burst(now, self.last_audio_at): + self._finish(reason=reason) + + def observe( + self, + *, + rms: int, + dbfs: float, + frame: int, + now: float, + start_fields: Optional[Dict[str, Any]] = None, + ) -> Optional[int]: + self.maybe_end(now=now, reason="silence_gap") + if rms <= RMS_TH: + return None + + new_burst = not self.active + if new_burst: + self.seq += 1 + self.active = True + self.first_frame = frame + self.active_frames = 0 + self.peak_rms = 0 + self.peak_dbfs = -120.0 + + self.last_frame = frame + self.last_audio_at = now + self.active_frames += 1 + if rms >= self.peak_rms: + self.peak_rms = rms + self.peak_dbfs = dbfs + + if new_burst: + log_flow_event( + logger, + self.start_step, + seq=self.seq, + rms=rms, + dbfs=f"{dbfs:.2f}", + frame=frame, + **(start_fields or {}), + ) + return self.seq + + return None + + def _finish(self, *, reason: str) -> None: + if FLOW_LOG_AUDIO_BURST_END: + duration_ms = max(FRAME_MS, (self.last_frame - self.first_frame + 1) * FRAME_MS) + log_flow_event( + logger, + self.end_step, + seq=self.seq, + reason=reason, + duration_ms=duration_ms, + active_frames=self.active_frames, + first_frame=self.first_frame, + last_frame=self.last_frame, + peak_rms=self.peak_rms, + peak_dbfs=f"{self.peak_dbfs:.2f}", + ) + self.active = False + + +class BridgeDebugState: + """ + Debug timeline (somente DEBUG). + Mantive pra você poder ligar nível DEBUG quando precisar investigar latência, + sem poluir INFO. + """ + + def __init__(self, *, base_t: float, timeline: Optional[CallTimeline] = None): + self.base_t = base_t + self.timeline = timeline + self.first_ws_in: Optional[float] = None + self.first_lk_out: Optional[float] = None + self.in_frames: int = 0 + self.out_frames: int = 0 + + def tlog(self, msg: str, **kv): + logger.debug("[T+%.3fs] %s | %s", time.monotonic() - self.base_t, msg, kv) + + def mark_ws_in(self): + if self.first_ws_in is None: + self.first_ws_in = time.monotonic() + self.tlog("WS->BRIDGE first_frame", in_frames=self.in_frames) + if self.timeline is not None: + self.timeline.emit( + "client_audio_first_frame_received", + in_frames=self.in_frames, + ) + + def mark_lk_out(self, backlog: int): + if self.first_lk_out is None: + self.first_lk_out = time.monotonic() + self.tlog("BRIDGE->LK first_capture", out_frames=self.out_frames, backlog=backlog) + if self.timeline is not None: + self.timeline.emit( + "client_audio_first_frame_published", + out_frames=self.out_frames, + backlog_frames=backlog, + ) + + +def make_token(identity: str, room_name: str) -> str: + grants = VideoGrants(room_join=True, room=room_name) + return ( + AccessToken( + api_key=os.environ["LIVEKIT_API_KEY"], + api_secret=os.environ["LIVEKIT_API_SECRET"], + ) + .with_identity(identity) + .with_grants(grants) + .to_jwt() + ) + + +def is_agent(p: rtc.Participant) -> bool: + try: + return int(p.kind) == 4 + except Exception: + ident = str(getattr(p, "identity", "") or "") + return ident == "agent" or ident.startswith("agent") + + +@app.get("/voice-client") +def voice_client() -> FileResponse: + return FileResponse(VOICE_CLIENT_HTML) + + +@app.websocket("/fake-agent/ws") +async def fake_remote_agent_ws(ws: WebSocket): + await ws.accept() + logger.info("FAKE_REMOTE_AGENT_CONNECT") + + try: + while True: + raw = await ws.receive_text() + payload = json.loads(raw) + response = build_fake_remote_agent_response(payload) + logger.info( + "FAKE_REMOTE_AGENT_REPLY | agent=%s | stage=%s", + str(response.get("agent") or payload.get("agent") or payload.get("action") or "-"), + str(response.get("stage") or "-"), + ) + await ws.send_text(json.dumps(response, ensure_ascii=False)) + except WebSocketDisconnect: + logger.info("FAKE_REMOTE_AGENT_DISCONNECT") + except Exception: + logger.exception("FAKE_REMOTE_AGENT_FAIL") + try: + await ws.send_text(json.dumps({"error": "fake_remote_agent_failed"})) + except Exception: + pass + try: + await close_websocket(ws, code=1011, reason="fake_agent_failed") + except Exception: + pass + + +async def dispatch_agent(lk_url: str, room_name: str, metadata: Dict[str, Any]) -> None: + lk = lkapi.LiveKitAPI( + url=lk_url, + api_key=os.environ["LIVEKIT_API_KEY"], + api_secret=os.environ["LIVEKIT_API_SECRET"], + ) + try: + meta_str = json.dumps(metadata, ensure_ascii=False) + await lk.agent_dispatch.create_dispatch( + lkapi.CreateAgentDispatchRequest( + agent_name=AGENT_NAME, + room=room_name, + metadata=meta_str, + ) + ) + finally: + await lk.aclose() + +async def ws_audio_receiver( + ws: WebSocket, + out_q: asyncio.Queue[bytes], + dbg: Optional[BridgeDebugState] = None, + *, + audio_enabled: Optional[asyncio.Event] = None, + recording_sink: Any = None, + input_backlog_config: Optional[AudioInputBacklogConfig] = None, + input_latency_tracker: Optional[AudioInputLatencyTracker] = None, +) -> None: + """ + Pipe puro: + - sempre lê do WS (evita backpressure/buffer) + - só enfileira frames quando audio_enabled estiver setado + - não faz lógica de "barge-in" aqui + """ + buf = bytearray() + frame_idx = 0 + + enabled_prev = False + dropped_bytes = 0 + client_bursts = FlowAudioBurstLogger( + start_step="audio_in", + end_step="audio_in_end", + ) + + while True: + msg = await ws.receive() + + if msg.get("type") == "websocket.disconnect": + raise WebSocketDisconnect(code=msg.get("code", 1000)) + + data = msg.get("bytes") + if not data: + continue + + enabled = True if audio_enabled is None else audio_enabled.is_set() + + # enquanto não estiver liberado, consome e descarta + if not enabled: + dropped_bytes += len(data) + buf.clear() + continue + + # transição: bridge liberou captura do cliente + if enabled and not enabled_prev: + enabled_prev = True + buf.clear() + frame_idx = 0 + now_enabled = time.monotonic() + if input_latency_tracker is not None: + input_latency_tracker.mark_enabled(now_enabled) + input_latency_tracker.maybe_log_latency( + frame_count=0, + now=now_enabled, + queue_size_frames=out_q.qsize(), + reason="receiver_enabled", + force=True, + ) + if dropped_bytes: + logger.info("CLIENT_AUDIO_ENABLED | dropped_bytes_before_enable=%s", dropped_bytes) + + buf.extend(data) + while len(buf) >= BYTES_PER_FRAME: + frame_bytes = bytes(buf[:BYTES_PER_FRAME]) + del buf[:BYTES_PER_FRAME] + + if dbg is not None: + dbg.in_frames += 1 + if dbg.in_frames == 1: + dbg.mark_ws_in() + frame_count = dbg.in_frames if dbg is not None else frame_idx + 1 + + if recording_sink is not None: + try: + recording_sink.record_client_frame(frame_bytes) + except Exception: + pass + + # mantém apenas atividade / medição (não interfere em áudio do agente) + now = time.monotonic() + if frame_idx % SILENCE_CHECK_EVERY == 0: + try: + rms = audioop.rms(frame_bytes, 2) + dbfs = pcm16_dbfs(frame_bytes) if rms > RMS_TH else -120.0 + seq = client_bursts.observe( + rms=rms, + dbfs=dbfs, + frame=frame_count, + now=now, + start_fields=_audio_input_latency_burst_fields( + input_latency_tracker, + frame_count=frame_count, + now=now, + queue_size_frames=out_q.qsize(), + ), + ) + if seq is not None: + if dbg is not None and dbg.timeline is not None: + dbg.timeline.emit( + "client_audio_activity_detected", + seq=seq, + rms=rms, + dbfs=round(dbfs, 2), + ) + elif rms <= RMS_TH: + client_bursts.maybe_end(now=now, reason="silence_gap") + except Exception: + pass + frame_idx += 1 + + try: + out_q.put_nowait(frame_bytes) + except asyncio.QueueFull: + queue_before = out_q.qsize() + selective = ShedResult() + if input_backlog_config is not None and input_backlog_config.energy_shed_enabled: + selective = shed_queue_backlog( + out_q, + keep_frames=max(0, queue_before - 1), + rms_threshold=input_backlog_config.silence_rms_threshold, + blind=False, + ) + if selective.dropped and input_latency_tracker is not None: + input_latency_tracker.record_shed( + source="receiver_queue_full", + dropped_frames=selective.dropped, + mode=selective.mode, + dropped_silent=selective.dropped_silent, + dropped_voiced=selective.dropped_voiced, + queue_before_frames=queue_before, + queue_after_frames=out_q.qsize(), + frame_count=frame_count, + now=now, + ) + try: + out_q.put_nowait(frame_bytes) + except asyncio.QueueFull: + try: + _ = out_q.get_nowait() + except asyncio.QueueEmpty: + pass + else: + if input_latency_tracker is not None: + input_latency_tracker.record_shed( + source="receiver_queue_full", + dropped_frames=1, + mode="queue_full_blind", + dropped_voiced=1, + queue_before_frames=queue_before, + queue_after_frames=max(0, out_q.qsize() - 1), + frame_count=frame_count, + now=now, + ) + out_q.put_nowait(frame_bytes) + + if input_backlog_config is not None: + _shed_audio_queue_backlog( + out_q, + config=input_backlog_config, + tracker=input_latency_tracker, + source="receiver", + frame_count=frame_count, + now=now, + ) + + if input_latency_tracker is not None: + input_latency_tracker.maybe_log_latency( + frame_count=frame_count, + now=now, + queue_size_frames=out_q.qsize(), + reason="receiver_audio", + ) + + +async def connect_publish_livekit( + room: rtc.Room, + lk_url: str, + token: str, + track: rtc.LocalAudioTrack, + opts: rtc.TrackPublishOptions, + track_published: asyncio.Event, + track_subscribed: asyncio.Event, + dbg: Optional[BridgeDebugState] = None, + timeline: Optional[CallTimeline] = None, +) -> None: + try: + await room.connect(lk_url, token) + if timeline is not None: + timeline.emit("livekit_room_connected") + publication = await room.local_participant.publish_track(track, opts) + + track_published.set() + if dbg: + dbg.tlog("LK connect+publish OK") + if timeline is not None: + timeline.emit("livekit_track_published") + + await publication.wait_for_subscription() + track_subscribed.set() + if dbg: + dbg.tlog("LK track subscribed") + if timeline is not None: + timeline.emit("livekit_track_subscribed") + except Exception: + logger.exception("LIVEKIT_SETUP_FAIL | connect/publish") + if timeline is not None: + timeline.emit("livekit_setup_failed") + raise + + +async def publish_queue_to_livekit( + source: rtc.AudioSource, + in_q: asyncio.Queue[bytes], + track_subscribed: asyncio.Event, + dbg: Optional[BridgeDebugState] = None, + *, + audio_enabled: Optional[asyncio.Event] = None, + input_backlog_config: Optional[AudioInputBacklogConfig] = None, + input_latency_tracker: Optional[AudioInputLatencyTracker] = None, +) -> None: + if dbg: + dbg.tlog("BRIDGE waiting track_subscribed") + + await track_subscribed.wait() + + # só começa a publicar áudio do cliente depois da liberação do bridge + if audio_enabled is not None: + if dbg: + dbg.tlog("BRIDGE waiting audio_enabled") + await audio_enabled.wait() + + now = time.monotonic() + if input_latency_tracker is not None: + input_latency_tracker.mark_enabled(now) + + frame_count = dbg.in_frames if dbg is not None else 0 + if input_backlog_config is not None: + _shed_audio_queue_backlog( + in_q, + config=input_backlog_config, + tracker=input_latency_tracker, + source="publisher_start", + frame_count=frame_count, + now=now, + ) + if input_latency_tracker is not None: + input_latency_tracker.maybe_log_latency( + frame_count=frame_count, + now=now, + queue_size_frames=in_q.qsize(), + reason="publisher_start", + force=True, + ) + + backlog0 = in_q.qsize() + publish_frame_idx = 0 + publish_bursts = FlowAudioBurstLogger( + start_step="audio_to_livekit", + end_step="audio_to_livekit_end", + ) + if dbg: + backlog_frames = in_q.qsize() + backlog_s = backlog_frames * (FRAME_MS / 1000.0) + dbg.tlog("BRIDGE start publish", backlog_frames=backlog_frames, backlog_s=round(backlog_s, 3)) + + if LK_CATCHUP_ENABLED and backlog0 > 0 and not ( + input_backlog_config is not None and input_backlog_config.shed_enabled + ): + keep_frames = max(1, int(LK_CATCHUP_KEEP_MS / FRAME_MS)) + if backlog0 > keep_frames: + if dbg: + dbg.tlog("BRIDGE catchup begin", backlog=backlog0, keep_frames=keep_frames, keep_ms=LK_CATCHUP_KEEP_MS) + + drained = 0 + while in_q.qsize() > keep_frames and drained < LK_CATCHUP_MAX_FAST_FRAMES: + try: + frame_bytes = in_q.get_nowait() + except asyncio.QueueEmpty: + break + + if dbg is not None: + dbg.out_frames += 1 + if dbg.out_frames == 1: + dbg.mark_lk_out(in_q.qsize()) + + frame = rtc.AudioFrame.create( + sample_rate=SAMPLE_RATE, + num_channels=CHANNELS, + samples_per_channel=SAMPLES_PER_FRAME, + ) + mv = frame.data.cast("B") + mv[:BYTES_PER_FRAME] = frame_bytes + + await source.capture_frame(frame) + drained += 1 + + if dbg: + dbg.tlog("BRIDGE catchup end", drained=drained, backlog_now=in_q.qsize()) + + next_t = time.monotonic() + while True: + now = time.monotonic() + frame_count = dbg.in_frames if dbg is not None else publish_frame_idx + if input_backlog_config is not None: + shed_result = _shed_audio_queue_backlog( + in_q, + config=input_backlog_config, + tracker=input_latency_tracker, + source="publisher_loop", + frame_count=frame_count, + now=now, + ) + if shed_result.dropped and shed_result.mode == "blind": + _maybe_clear_livekit_source_queue( + source, + tracker=input_latency_tracker, + reason="publisher_loop_shed", + frame_count=frame_count, + queue_size_frames=in_q.qsize(), + now=now, + ) + next_t = time.monotonic() + + frame_bytes = await in_q.get() + + if dbg is not None: + dbg.out_frames += 1 + if dbg.out_frames == 1: + dbg.mark_lk_out(in_q.qsize()) + + now = time.monotonic() + if publish_frame_idx % SILENCE_CHECK_EVERY == 0: + try: + rms = audioop.rms(frame_bytes, 2) + dbfs = pcm16_dbfs(frame_bytes) if rms > RMS_TH else -120.0 + if rms > RMS_TH: + publish_bursts.observe( + rms=rms, + dbfs=dbfs, + frame=dbg.out_frames if dbg is not None else publish_frame_idx, + now=now, + start_fields={ + "backlog_frames": in_q.qsize(), + "livekit_queued_ms": max(0, round(source.queued_duration * 1000)), + **_audio_input_latency_burst_fields( + input_latency_tracker, + frame_count=dbg.in_frames if dbg is not None else publish_frame_idx + 1, + now=now, + queue_size_frames=in_q.qsize(), + ), + }, + ) + else: + publish_bursts.maybe_end(now=now, reason="silence_gap") + except Exception: + pass + publish_frame_idx += 1 + + if now - next_t > 0.2: + next_t = now + + if now < next_t: + await asyncio.sleep(next_t - now) + + frame = rtc.AudioFrame.create( + sample_rate=SAMPLE_RATE, + num_channels=CHANNELS, + samples_per_channel=SAMPLES_PER_FRAME, + ) + mv = frame.data.cast("B") + mv[:BYTES_PER_FRAME] = frame_bytes + + await source.capture_frame(frame) + next_t += (FRAME_MS / 1000.0) + + +async def stream_agent_audio_to_queue( + agent_q: asyncio.Queue[bytes], + agent_participant: rtc.Participant, + activity: AudioActivity, + timeline: Optional[CallTimeline] = None, + call_context: Optional[Mapping[str, Any]] = None, +): + flow_context = dict(call_context or {}) + stream = rtc.AudioStream.from_participant( + participant=agent_participant, + track_source=rtc.TrackSource.SOURCE_MICROPHONE, + sample_rate=SAMPLE_RATE, + num_channels=CHANNELS, + frame_size_ms=FRAME_MS, + ) + + buf = bytearray() + frame_idx = 0 + + first_frame_sent = False + last_agent_audio_at = 0.0 + agent_audio_burst_seq = 0 + dropped_queue_frames = 0 + + async for ev in stream: + b = bytes(ev.frame.data) + buf.extend(b) + + while len(buf) >= BYTES_PER_FRAME: + frame = bytes(buf[:BYTES_PER_FRAME]) + del buf[:BYTES_PER_FRAME] + + now = time.monotonic() + + # mede RMS só para telemetria/atividade — nunca para descartar frame + rms = None + if frame_idx % AGENT_OUT_RMS_EVERY == 0: + try: + rms = audioop.rms(frame, 2) + except Exception: + rms = None + frame_idx += 1 + + # todo frame do agente é encaminhado bit a bit; a energia só marca + # "fala ativa" (usado por idle/hold e mock-stop), não filtra o áudio. + if rms is not None and rms > AGENT_OUT_SPEECH_RMS_TH: + new_burst = _is_new_audio_burst(now, last_agent_audio_at) + last_agent_audio_at = now + activity.last_agent_out = now + if new_burst: + agent_audio_burst_seq += 1 + log_flow_event( + logger, + "audio_from_agent", + **flow_context, + seq=agent_audio_burst_seq, + participant=str(getattr(agent_participant, "identity", "") or ""), + rms=rms if rms is not None else "", + frame=frame_idx, + ) + if timeline is not None: + timeline.emit( + "agent_audio_activity_detected", + **flow_context, + seq=agent_audio_burst_seq, + participant_identity=str(getattr(agent_participant, "identity", "") or ""), + rms=rms, + ) + if not first_frame_sent: + first_frame_sent = True + if timeline is not None: + timeline.emit( + "agent_audio_first_frame_received", + **flow_context, + participant_identity=str(getattr(agent_participant, "identity", "") or ""), + ) + + try: + agent_q.put_nowait(frame) + except asyncio.QueueFull: + dropped_queue_frames += 1 + if dropped_queue_frames == 1 or dropped_queue_frames % 50 == 0: + log_flow_event( + logger, + "audio_from_agent_queue_overflow", + **flow_context, + participant=str(getattr(agent_participant, "identity", "") or ""), + dropped_agent_audio_frames=dropped_queue_frames, + dropped_voiced_frames=1, + queue_size=agent_q.qsize(), + content_policy="oldest_frame", + sync_action="drop_audio_to_catch_up", + ) + if timeline is not None: + timeline.emit( + "agent_audio_queue_overflow", + **flow_context, + participant_identity=str(getattr(agent_participant, "identity", "") or ""), + dropped_agent_audio_frames=dropped_queue_frames, + dropped_voiced_frames=1, + queue_size=agent_q.qsize(), + content_policy="oldest_frame", + sync_action="drop_audio_to_catch_up", + ) + _ = agent_q.get_nowait() + agent_q.put_nowait(frame) + + +@app.on_event("startup") +async def _startup(): + _load_runtime_env() + + +@app.websocket("/ws/agent") +async def ws_agent(ws: WebSocket): + _load_runtime_env() + call_started_unix_ms = round(time.time() * 1000) + await ws.accept() + + lk_url = os.environ["LIVEKIT_URL"] + call_t0 = time.monotonic() + capacity_lease = None + ws_close_code = 1000 + ws_close_reason = "ws_agent_cleanup" + ws_client_disconnected = False + bridge_output_stats: Optional[BridgeOutputStats] = None + bridge_log_context: Dict[str, Any] = {} + entire_call_recorder = None + readiness_task: asyncio.Task[Any] | None = None + try: + start_ctx = await recv_start_message(ws, logger=logger) + reservation = await READINESS_GATE.reserve_capacity() + if not reservation.allowed: + logger.warning( + "WS_AGENT_PRE_READY_STOP | reason=capacity_exceeded | active_connections=%s | max_connections=%s", + reservation.active_connections, + reservation.max_connections, + ) + await ws.send_text(json.dumps(build_capacity_stop_message(), ensure_ascii=False)) + try: + await close_websocket(ws, code=1013, reason="pre_ready:capacity_exceeded") + except Exception: + pass + return + capacity_lease = reservation.lease + + agent_name = str(start_ctx.data.get("agent") or "").strip().lower() or "conta" + readiness_task = asyncio.create_task( + READINESS_GATE.evaluate_resources(agent_name=agent_name, call_config=start_ctx.call_config), + name="evaluate_call_readiness", + ) + + bootstrap = build_bridge_session_bootstrap( + start_ctx=start_ctx, + app_env=APP_ENV, + livekit_room=os.getenv("LIVEKIT_ROOM", ""), + default_protocol=os.getenv("DEFAULT_PROTOCOL", ""), + token_factory=make_token, + logger=logger, + timeline_origin_unix_ms=call_started_unix_ms, + ) + call_id_ged = bootstrap.call_id_ged + call_config = bootstrap.call_config + remote_agent_context = bootstrap.remote_agent_context + room_name = bootstrap.room_name + identity = bootstrap.identity + token = bootstrap.token + protocol = bootstrap.protocol + session_id = str(bootstrap.dispatch_metadata.get("session_id") or "") + phone_number = bootstrap.phone_number + timeline = bootstrap.timeline + + logger.info( + "CALL_START | room=%s | protocol=%s | session_id=%s | call_id_ged=%s | identity=%s", + room_name, + protocol, + session_id or "-", + call_id_ged or "-", + identity, + ) + log_flow_event( + logger, + "call_start", + room=room_name, + protocol=protocol, + session_id=session_id, + call_id_ged=call_id_ged, + phone_number=phone_number, + bridge=identity, + ) + timeline.emit( + "call_start", + call_id_ged=call_id_ged, + bridge_identity=identity, + remote_agent=remote_agent_context.get("agent", ""), + agent_backend=resolve_agent_backend_name(call_config, os.getenv("AGENT_BACKEND", "remote_ws")), + stt_provider=( + resolve_stt_overrides(call_config).get("provider") + or os.getenv("STT_PROVIDER", "internal_http") + ), + tts_provider=( + resolve_tts_overrides(call_config).get("provider") + or os.getenv("TTS_PROVIDER", "elevenlabs") + ), + timeline_file=str(timeline.path), + ) + timeline.emit("readiness_started", mode="parallel_post_ready") + timeline.emit("bridge_start_received") + + + # READY imediato + await ws.send_text( + json.dumps( + { + "type": "ready", + "room": room_name, + "session_id": bootstrap.dispatch_metadata.get("session_id", ""), + "sample_rate": SAMPLE_RATE, + "channels": CHANNELS, + "frame_ms": FRAME_MS, + "bytes_per_frame": BYTES_PER_FRAME, + "debug_events_enabled": start_ctx.debug_events_enabled, + "stress_test": start_ctx.stress_test, + } + ) + ) + timeline.emit("ready_sent") + + now = time.monotonic() + activity = AudioActivity(last_agent_out=now) + + client_q: asyncio.Queue[bytes] = asyncio.Queue(maxsize=1500) # ~30s + agent_q: asyncio.Queue[bytes] = asyncio.Queue(maxsize=600) # mais folga pra áudio do agente + + dbg = BridgeDebugState(base_t=time.monotonic(), timeline=timeline) + bridge_output_stats = BridgeOutputStats() + bridge_log_context = { + "room": room_name, + "protocol": protocol, + "session_id": session_id, + "call_id_ged": call_id_ged, + "phone_number": phone_number, + "bridge": identity, + } + lifecycle = RoomLifecycleState() + + def _publish_bridge_debug_event( + event: str, payload: Mapping[str, Any] + ) -> None: + if not start_ctx.debug_events_enabled: + return + + async def _send() -> None: + try: + await ws.send_text( + json.dumps( + { + "type": "debug_event", + "stress_test": True, + "event": event, + "data": dict(payload), + }, + ensure_ascii=False, + ) + ) + except Exception: + logger.debug("bridge debug event forwarding failed", exc_info=True) + + task = asyncio.create_task( + _send(), name=f"forward_{event.replace('.', '_')}" + ) + lifecycle.debug_forward_tasks.add(task) + task.add_done_callback(lifecycle.debug_forward_tasks.discard) + + input_backlog_config = audio_input_backlog_config_from_env(call_config) + input_latency_tracker = AudioInputLatencyTracker( + input_backlog_config, + timeline=timeline, + call_context=bridge_log_context, + debug_event_publisher=_publish_bridge_debug_event, + ) + input_backlog_config_payload = { + **bridge_log_context, + "shed_enabled": input_backlog_config.shed_enabled, + "shed_threshold_ms": input_backlog_config.shed_threshold_ms, + "shed_keep_ms": input_backlog_config.shed_keep_ms, + "latency_metrics_enabled": input_backlog_config.latency_metrics_enabled, + "latency_alert_ms": input_backlog_config.latency_alert_ms, + "latency_log_interval_s": input_backlog_config.latency_log_interval_s, + "livekit_source_queue_size_ms": input_backlog_config.livekit_source_queue_size_ms, + "livekit_source_clear_on_shed": input_backlog_config.livekit_source_clear_on_shed, + "config_source": input_backlog_config.config_source, + "call_config_overrides": list(input_backlog_config.call_config_overrides), + } + log_flow_event(logger, "audio_in_backlog_config", **input_backlog_config_payload) + timeline.emit("audio_in_backlog_config", **input_backlog_config_payload) + entire_call_recorder = create_entire_call_recorder_from_env( + session_id=session_id, + sample_rate=SAMPLE_RATE, + channels=CHANNELS, + sample_width=2, + frame_ms=FRAME_MS, + bytes_per_frame=BYTES_PER_FRAME, + logger_override=logger, + timeline=timeline, + metadata=bridge_log_context, + ) + first_agent_audio_sent = asyncio.Event() + mock_stop_after_first_audio_enabled = ( + os.getenv("MOCK_STOP_AFTER_FIRST_AUDIO_ENABLED", "0").strip().lower() + in {"1", "true", "yes", "on"} + ) + try: + mock_stop_after_first_audio_silence_s = max( + 0.0, + float(os.getenv("MOCK_STOP_AFTER_FIRST_AUDIO_SILENCE_S", "0.35")), + ) + except ValueError: + mock_stop_after_first_audio_silence_s = 0.35 + mock_stop_after_first_audio_reason = ( + os.getenv("MOCK_STOP_AFTER_FIRST_AUDIO_REASON", "DONE").strip() + or "DONE" + ) + + # gate: o bridge só aceita áudio do cliente depois de liberar a sessão + client_audio_enabled = asyncio.Event() # começa bloqueado + + # receiver sempre lendo (pra não estourar), mas só enfileira depois do gate + t_rx = asyncio.create_task( + ws_audio_receiver( + ws, + client_q, + dbg=dbg, + audio_enabled=client_audio_enabled, + recording_sink=entire_call_recorder, + input_backlog_config=input_backlog_config, + input_latency_tracker=input_latency_tracker, + ), + name="ws_audio_receiver", + ) + + room = rtc.Room() + agent_ready = asyncio.Event() + track_published = asyncio.Event() + track_subscribed = asyncio.Event() + + call_done = asyncio.Event() + agent_reconnect_enabled = _env_bool("AGENT_RECONNECT_ENABLED", True) + agent_reconnect_max_attempts = max(0, _env_int("AGENT_RECONNECT_MAX_ATTEMPTS", 1)) + agent_reconnect_timeout_s = max(0.1, _env_float("AGENT_RECONNECT_TIMEOUT_S", 10.0)) + agent_reconnect_poll_s = 0.05 + agent_reconnect_lock = asyncio.Lock() + + async def _dispatch_agent_with_timeline(*, reason: str, attempt: int = 0) -> None: + fields: Dict[str, Any] = {"agent_name": AGENT_NAME, "reason": reason} + if attempt: + fields["attempt"] = attempt + timeline.emit("dispatch_started", **fields) + log_flow_event( + logger, + "agent_dispatch", + agent_name=AGENT_NAME, + room=room_name, + reason=reason, + attempt=attempt, + ) + try: + await dispatch_agent(lk_url, room_name, bootstrap.dispatch_metadata) + except Exception: + timeline.emit("dispatch_failed", **fields) + raise + timeline.emit("dispatch_completed", **fields) + + async def _wait_for_agent_reconnect(previous_generation: int) -> bool: + while not call_done.is_set(): + if ( + lifecycle.agent_participant is not None + and lifecycle.agent_generation > previous_generation + ): + return True + await asyncio.sleep(agent_reconnect_poll_s) + return False + + def _stop_for_agent_disconnect(agent_identity: str) -> None: + if call_done.is_set(): + return + lifecycle.done_payload = { + "stage": "DONE", + "status": "stop_agent_runtime_unavailable", + "reason": "agent_disconnected", + "resource": "agent_runtime", + "failed_resources": ["agent_runtime"], + "phase": "in_session", + } + logger.error( + "AGENT_RECONNECT_EXHAUSTED | room=%s | protocol=%s | agent_identity=%s", + room_name, + protocol, + agent_identity or "-", + ) + timeline.emit( + "agent_reconnect_exhausted", + agent_identity=agent_identity or "", + max_attempts=agent_reconnect_max_attempts, + timeout_ms=round(agent_reconnect_timeout_s * 1000), + ) + call_done.set() + + async def _recover_agent_disconnect(agent_identity: str, generation: int) -> None: + if call_done.is_set(): + return + async with agent_reconnect_lock: + if call_done.is_set(): + return + if lifecycle.agent_participant is not None and lifecycle.agent_generation > generation: + return + if not agent_reconnect_enabled or agent_reconnect_max_attempts <= 0: + _stop_for_agent_disconnect(agent_identity) + return + + for attempt in range(1, agent_reconnect_max_attempts + 1): + if call_done.is_set(): + return + logger.warning( + "AGENT_RECONNECT_ATTEMPT | room=%s | protocol=%s | agent_identity=%s | attempt=%s/%s | timeout_s=%.3f", + room_name, + protocol, + agent_identity or "-", + attempt, + agent_reconnect_max_attempts, + agent_reconnect_timeout_s, + ) + timeline.emit( + "agent_reconnect_attempt", + agent_identity=agent_identity or "", + attempt=attempt, + max_attempts=agent_reconnect_max_attempts, + timeout_ms=round(agent_reconnect_timeout_s * 1000), + ) + try: + await _dispatch_agent_with_timeline(reason="agent_reconnect", attempt=attempt) + except Exception: + logger.exception( + "AGENT_RECONNECT_DISPATCH_FAIL | room=%s | protocol=%s | attempt=%s/%s", + room_name, + protocol, + attempt, + agent_reconnect_max_attempts, + ) + timeline.emit( + "agent_reconnect_dispatch_failed", + agent_identity=agent_identity or "", + attempt=attempt, + max_attempts=agent_reconnect_max_attempts, + ) + + try: + reconnected = await asyncio.wait_for( + _wait_for_agent_reconnect(generation), + timeout=agent_reconnect_timeout_s, + ) + except asyncio.TimeoutError: + reconnected = False + + if call_done.is_set(): + return + + if reconnected: + new_identity = str(getattr(lifecycle.agent_participant, "identity", "") or "") + logger.info( + "AGENT_RECONNECT_OK | room=%s | protocol=%s | old_agent_identity=%s | new_agent_identity=%s | attempt=%s", + room_name, + protocol, + agent_identity or "-", + new_identity or "-", + attempt, + ) + timeline.emit( + "agent_reconnect_completed", + old_agent_identity=agent_identity or "", + new_agent_identity=new_identity, + attempt=attempt, + ) + return + + timeline.emit( + "agent_reconnect_timeout", + agent_identity=agent_identity or "", + attempt=attempt, + max_attempts=agent_reconnect_max_attempts, + ) + + _stop_for_agent_disconnect(agent_identity) + + t_dispatch0 = time.monotonic() + register_room_lifecycle_handlers( + room=room, + logger=logger, + timeline=timeline, + room_name=room_name, + protocol=protocol, + dispatch_started_at=t_dispatch0, + agent_ready=agent_ready, + call_done=call_done, + state=lifecycle, + is_agent=is_agent, + on_agent_disconnect=_recover_agent_disconnect, + ws=ws, + debug_events_enabled=start_ctx.debug_events_enabled, + ) + t_done = asyncio.create_task( + watch_call_done( + call_done=call_done, + state=lifecycle, + ws=ws, + logger=logger, + room_name=room_name, + protocol=protocol, + timeline=timeline, + ), + name="watch_call_done", + ) + + async def _watch_parallel_readiness() -> None: + try: + readiness = await readiness_task + except asyncio.CancelledError: + raise + except Exception as exc: + logger.exception("WS_AGENT_READINESS_FAIL | agent=%s", agent_name) + timeline.emit("readiness_failed", error=type(exc).__name__) + if not call_done.is_set(): + lifecycle.done_payload = build_bridge_failed_stop_message()["data"] + call_done.set() + return + + if readiness.is_healthy(): + timeline.emit( + "readiness_ok", + cached=readiness.cached, + active_connections=readiness.active_connections, + max_connections=readiness.max_connections, + ) + return + + logger.warning( + "WS_AGENT_READINESS_FAIL | agent=%s | failed_resources=%s | report=%s", + agent_name, ",".join(readiness.failed_resources), json.dumps(readiness.as_dict(), ensure_ascii=False), + ) + timeline.emit("readiness_failed", failed_resources=readiness.failed_resources) + if not call_done.is_set(): + stop_message = readiness.build_stop_message() + stop_message["data"]["phase"] = "in_session" + lifecycle.done_payload = stop_message["data"] + call_done.set() + + t_readiness = asyncio.create_task(_watch_parallel_readiness(), name="watch_call_readiness") + + try: + source = rtc.AudioSource( + SAMPLE_RATE, + CHANNELS, + queue_size_ms=input_backlog_config.livekit_source_queue_size_ms, + ) + except TypeError: + source = rtc.AudioSource(SAMPLE_RATE, CHANNELS) + livekit_source_config_payload = { + **bridge_log_context, + "queue_size_ms": input_backlog_config.livekit_source_queue_size_ms, + "clear_on_shed": input_backlog_config.livekit_source_clear_on_shed, + } + log_flow_event(logger, "livekit_audio_source_config", **livekit_source_config_payload) + timeline.emit("livekit_audio_source_config", **livekit_source_config_payload) + + track = rtc.LocalAudioTrack.create_audio_track("ws-mic", source) + opts = rtc.TrackPublishOptions() + opts.source = rtc.TrackSource.SOURCE_MICROPHONE + + # dispatch do agent + t_dispatch = asyncio.create_task( + _dispatch_agent_with_timeline(reason="initial"), + name="dispatch_agent", + ) + + # conecta + publica track + espera subscribe + t_lk = asyncio.create_task( + connect_publish_livekit( + room=room, + lk_url=lk_url, + token=token, + track=track, + opts=opts, + track_published=track_published, + track_subscribed=track_subscribed, + dbg=dbg, + timeline=timeline, + ), + name="connect_publish_livekit", + ) + + # publish do áudio do cliente -> LK (só depois do gate) + t_in = asyncio.create_task( + publish_queue_to_livekit( + source, + client_q, + track_subscribed, + dbg=dbg, + audio_enabled=client_audio_enabled, + input_backlog_config=input_backlog_config, + input_latency_tracker=input_latency_tracker, + ), + name="publish_queue_to_livekit", + ) + + input_latency_tracker.mark_enabled(time.monotonic()) + await enable_client_audio( + logger=logger, + timeline=timeline, + client_audio_enabled=client_audio_enabled, + room_name=room_name, + protocol=protocol, + ) + t_control = asyncio.create_task( + notify_client_audio_enabled( + room=room, + logger=logger, + timeline=timeline, + room_name=room_name, + protocol=protocol, + track_published=track_published, + agent_ready=agent_ready, + lifecycle=lifecycle, + repeat=True, + ), + name="notify_client_audio_enabled", + ) + + ws_overrides = resolve_ws_overrides(call_config) + + # saída: agente -> cliente (via ws_out_loop) + t_out_loop = asyncio.create_task( + ws_out_loop( + ws, + agent_q, + frame_ms=FRAME_MS, + bytes_per_frame=BYTES_PER_FRAME, + first_agent_audio_sent=first_agent_audio_sent, + flow_logger=logger, + timeline=timeline, + call_context=bridge_log_context, + output_stats=bridge_output_stats, + output_gain=ws_overrides.get("output_gain"), + recording_sink=entire_call_recorder, + debug_event_publisher=_publish_bridge_debug_event, + ), + name="ws_out_loop", + ) + + t_agent_to_q = asyncio.create_task( + stream_agent_audio_when_ready( + agent_ready=agent_ready, + lifecycle=lifecycle, + agent_q=agent_q, + activity=activity, + timeline=timeline, + stream_agent_audio=( + lambda agent_q, agent_participant, activity, timeline=None: stream_agent_audio_to_queue( + agent_q, + agent_participant, + activity, + timeline=timeline, + call_context=bridge_log_context, + ) + ), + repeat=True, + ), + name="stream_agent_audio_when_ready", + ) + t_mock_stop = None + if mock_stop_after_first_audio_enabled: + logger.info( + "MOCK_STOP_AFTER_FIRST_AUDIO_ENABLED | room=%s | protocol=%s | silence_s=%.3f | reason=%s", + room_name, + protocol, + mock_stop_after_first_audio_silence_s, + mock_stop_after_first_audio_reason, + ) + timeline.emit( + "mock_stop_after_first_audio_enabled", + silence_ms=round(mock_stop_after_first_audio_silence_s * 1000), + reason=mock_stop_after_first_audio_reason, + ) + t_mock_stop = asyncio.create_task( + watch_mock_stop_after_first_audio( + ws=ws, + logger=logger, + timeline=timeline, + room_name=room_name, + protocol=protocol, + first_agent_audio_sent=first_agent_audio_sent, + activity=activity, + agent_q=agent_q, + silence_s=mock_stop_after_first_audio_silence_s, + reason=mock_stop_after_first_audio_reason, + ), + name="mock_stop_after_first_audio", + ) + + try: + tasks = {t_rx, t_out_loop, t_in, t_agent_to_q, t_lk, t_dispatch, t_done, t_readiness, t_control} + if t_mock_stop is not None: + tasks.add(t_mock_stop) + pending = set(tasks) + while pending: + done, pending = await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED) + if t_done in done: + break + for d in done: + exc = d.exception() + if exc: + logger.error( + "WS_AGENT_TASK_FAIL | room=%s | protocol=%s | task=%s", + room_name, + protocol, + d.get_name(), + exc_info=exc, + ) + raise exc + + except WebSocketDisconnect as exc: + ws_client_disconnected = True + logger.info( + "WS_DISCONNECT | room=%s | protocol=%s | code=%s", + room_name, + protocol, + getattr(exc, "code", ""), + ) + timeline.emit("ws_disconnect") + except Exception as exc: + ws_close_code = 1011 + ws_close_reason = "bridge_failed" + logger.exception("WS_AGENT_FAIL | room=%s | protocol=%s", room_name, protocol) + timeline.emit("bridge_failed") + try: + await ws.send_text(json.dumps(build_bridge_failed_stop_message(), ensure_ascii=False)) + except Exception: + pass + finally: + for t in (t_rx, t_out_loop, t_in, t_agent_to_q, t_lk, t_dispatch, t_done, t_readiness, t_control, t_mock_stop): + try: + t.cancel() + except Exception: + pass + + try: + await room.disconnect() + except Exception: + pass + + if not ws_client_disconnected: + try: + await close_websocket(ws, code=ws_close_code, reason=ws_close_reason) + except Exception: + pass + + if entire_call_recorder is not None: + try: + await entire_call_recorder.finalize() + except Exception: + logger.exception( + "ENTIRE_CALL_RECORDING_FINALIZE_FAIL | room=%s | protocol=%s | session_id=%s", + room_name, + protocol, + session_id or "-", + ) + + dur = time.monotonic() - call_t0 + audio_to_tia_summary: Dict[str, Any] = {} + if bridge_output_stats is not None: + audio_to_tia_summary = { + "destination": "ws_client", + "ws_close_reason": ws_close_reason, + "ws_client_disconnected": ws_client_disconnected, + "duration_ms": round(dur * 1000), + "total_ws_frames": bridge_output_stats.total_ws_frames, + "agent_audio_frames": bridge_output_stats.agent_audio_frames, + "silence_frames": bridge_output_stats.silence_frames, + "agent_audio_bursts": bridge_output_stats.agent_audio_bursts, + "first_agent_audio_sent": bridge_output_stats.first_agent_audio_sent, + "queued_agent_frames": agent_q.qsize(), + } + log_flow_event( + logger, + "audio_to_tia_summary", + **bridge_log_context, + **audio_to_tia_summary, + ) + timeline.emit( + "agent_audio_to_tia_summary", + **bridge_log_context, + **audio_to_tia_summary, + ) + logger.info( + "CALL_END | room=%s | protocol=%s | session_id=%s | duration_s=%.1f", + room_name, + protocol, + session_id or "-", + dur, + ) + timeline.emit( + "call_end", + duration_ms=round(dur * 1000), + client_in_frames=dbg.in_frames, + bridge_out_frames=dbg.out_frames, + audio_to_tia_total_ws_frames=audio_to_tia_summary.get("total_ws_frames", 0), + audio_to_tia_agent_audio_frames=audio_to_tia_summary.get("agent_audio_frames", 0), + audio_to_tia_agent_audio_bursts=audio_to_tia_summary.get("agent_audio_bursts", 0), + audio_to_tia_first_agent_audio_sent=audio_to_tia_summary.get("first_agent_audio_sent", False), + ) + finally: + if capacity_lease is not None: + await capacity_lease.release() + if readiness_task is not None and not readiness_task.done(): + readiness_task.cancel() + + +@app.websocket("/ws/text") +async def ws_text(ws: WebSocket): + from app.services.text_pipeline import TextSessionPipeline + + await ws.accept() + session_id = str(uuid.uuid4()) + t0 = time.monotonic() + logger.info("TEXT_START | session_id=%s", session_id) + try: + start_ctx = await recv_start_message(ws, logger=logger) + pipeline = TextSessionPipeline( + ws, + start_ctx.payload, + session_id, + session_data=start_ctx.session_data, + intro=start_ctx.intro, + ) + await pipeline.run() + except WebSocketDisconnect: + logger.info("TEXT_END | session_id=%s | reason=disconnect | duration_s=%.1f", session_id, time.monotonic() - t0) + except Exception as e: + logger.exception("TEXT_FAIL | session_id=%s", session_id) + try: + await ws.send_text(json.dumps({"type": "error", "message": str(e)})) + except Exception: + pass + try: + await close_websocket(ws, code=1011, reason="text_pipeline_failed") + except Exception: + pass + + +@app.websocket("/ws/text_stream") +async def ws_text_stream(ws: WebSocket): + from app.services.text_pipeline_stream import TextSessionPipelineStream + + await ws.accept() + session_id = str(uuid.uuid4()) + t0 = time.monotonic() + logger.info("TEXT_STREAM_START | session_id=%s", session_id) + try: + start_ctx = await recv_start_message(ws, logger=logger) + pipeline = TextSessionPipelineStream( + ws, + start_ctx.payload, + session_id, + session_data=start_ctx.session_data, + intro=start_ctx.intro, + ) + await pipeline.run() + except WebSocketDisconnect: + logger.info( + "TEXT_STREAM_END | session_id=%s | reason=disconnect | duration_s=%.1f", + session_id, + time.monotonic() - t0, + ) + except Exception as e: + logger.exception("TEXT_STREAM_FAIL | session_id=%s", session_id) + try: + await ws.send_text(json.dumps({"type": "error", "message": str(e)})) + except Exception: + pass + try: + await close_websocket(ws, code=1011, reason="text_stream_pipeline_failed") + except Exception: + pass + +@app.get("/health") +def health(): + return {"status": "ok"} + + +@app.get("/health/resources") +async def health_resources(): + report = await READINESS_GATE.evaluate_default_resources() + status_code = 200 if report.is_healthy() else 503 + return JSONResponse(status_code=status_code, content=report.as_dict()) + + +@app.get("/health/services") +async def health_services(): + report = await READINESS_GATE.evaluate_services() + status_code = 200 if report.is_healthy() else 503 + return JSONResponse(status_code=status_code, content=report.as_dict()) + +if __name__ == "__main__": + import uvicorn + import argparse + + parser = argparse.ArgumentParser(description="LiveKit WebSocket Gateway (PCM16 16k <-> LiveKit)") + parser.add_argument("--port", type=int, default=8000) + parser.add_argument("--host", type=str, default="0.0.0.0") + parser.add_argument("--log-level", type=str, default="info") + parser.add_argument("--reload", action="store_true") + args = parser.parse_args() + + uvicorn.run( + app, + host=args.host, + port=args.port, + ws_ping_interval=float(os.getenv("UVICORN_WS_PING_INTERVAL_S", "20")), + ws_ping_timeout=float(os.getenv("UVICORN_WS_PING_TIMEOUT_S", "20")), + ws_max_size=50 * 1024 * 1024, + log_level="info" if args.log_level is None else args.log_level, + reload=args.reload, + ) diff --git a/src/app/ws_gateway/readiness.py b/src/app/ws_gateway/readiness.py new file mode 100644 index 0000000..86a3513 --- /dev/null +++ b/src/app/ws_gateway/readiness.py @@ -0,0 +1,944 @@ +from __future__ import annotations + +import asyncio +import json +import os +import time +from dataclasses import dataclass +from typing import Any, Dict, Mapping, Optional +from urllib.parse import urlsplit, urlunsplit + +import aiohttp +import httpx + +from app.common.call_config import ( + normalize_call_config, + resolve_agent_backend_name, + resolve_stt_overrides, + resolve_tts_overrides, +) +from app.livekit.adapters.azure_rest_tts import AzureRESTTTS +from app.livekit.adapters.xai_tts import OraclexAITTS +from app.livekit.azure_speech import ( + AZURE_SPEECH_TTS_SAMPLE_RATE, + resolve_azure_speech_tts_config, +) +from app.providers.stt_fake import FakeSTT +from app.providers.tts import FakeTTS as FakeProviderTTS + + +STOP_STATUS_BY_RESOURCE = { + "capacity": "stop_capacity_tia", + "agent_runtime": "stop_agent_runtime_unavailable", + "agent_backend": "stop_agent_backend_unavailable", + "stt": "stop_stt_unavailable", + "tts": "stop_tts_unavailable", +} +STOP_RESOURCE_PRECEDENCE = ( + "capacity", + "agent_runtime", + "agent_backend", + "stt", + "tts", +) +HEALTHY_CHECK_STATUSES = {"ok", "skipped"} +PROBE_TEXT = "teste de prontidao" +XAI_TTS_READINESS_MODE_ENV = "XAI_TTS_READINESS_MODE" +XAI_TTS_READINESS_MODES = {"connect", "synthesize"} +DEFAULT_FINAL_STOP_STATUS_BY_REASON = { + "resolved": "stop_resolvido_e_finalizado", + "unresolved": "stop_nao_resolvido", + "other_subject": "stop_outro_assunto", + "long_silence": "stop_silencio_longo", +} +DEFAULT_FINAL_STOP_REASON_BY_KIND = { + "resolved": "stage_done", + "unresolved": "nao_resolvido", + "other_subject": "outro_assunto", + "long_silence": "no_user_response", +} + + +def build_stop_message( + *, + status: str, + reason: str, + resource: str = "", + failed_resources: Optional[list[str]] = None, + phase: str = "", +) -> Dict[str, Any]: + data: Dict[str, Any] = { + "status": str(status or "").strip(), + "reason": str(reason or "").strip(), + } + if resource: + data["resource"] = str(resource).strip() + if failed_resources: + data["failed_resources"] = [str(item).strip() for item in failed_resources if str(item).strip()] + if phase: + data["phase"] = str(phase).strip() + return {"type": "stop", "data": data} + + +def _final_stop_status_from_env(key: str) -> str: + env_name = f"FINAL_STOP_STATUS_{key.upper()}" + return ( + os.getenv(env_name, DEFAULT_FINAL_STOP_STATUS_BY_REASON[key]) or DEFAULT_FINAL_STOP_STATUS_BY_REASON[key] + ).strip() + + +def _format_probe_error(exc: Exception, *, resource: str, timeout_s: float) -> str: + """Return a useful, secret-free error summary for readiness failures.""" + error_type = type(exc).__name__ + detail = _strip_text(str(exc)) + if isinstance(exc, TimeoutError): + return f"{resource} readiness timed out after {round(timeout_s * 1000)}ms ({error_type})" + if detail: + return f"{resource} readiness failed ({error_type}): {detail}" + return f"{resource} readiness failed ({error_type})" + + +def _final_stop_reason_from_env(key: str) -> str: + env_name = f"FINAL_STOP_REASON_{key.upper()}" + return ( + os.getenv(env_name, DEFAULT_FINAL_STOP_REASON_BY_KIND[key]) or DEFAULT_FINAL_STOP_REASON_BY_KIND[key] + ).strip() + + +def resolve_completed_stop_status(reason: str) -> str: + normalized_reason = str(reason or "").strip().lower() + for key in DEFAULT_FINAL_STOP_STATUS_BY_REASON: + configured_reason = _final_stop_reason_from_env(key).lower() + if normalized_reason and normalized_reason == configured_reason: + return _final_stop_status_from_env(key) + + default_key = ( + os.getenv("FINAL_STOP_DEFAULT_KIND", "resolved") or "resolved" + ).strip().lower() + if default_key not in DEFAULT_FINAL_STOP_STATUS_BY_REASON: + default_key = "resolved" + return _final_stop_status_from_env(default_key) + + +def build_completed_stop_message(reason: str) -> Dict[str, Any]: + resolved_reason = str(reason or "").strip() or _final_stop_reason_from_env("resolved") + return build_stop_message( + status=resolve_completed_stop_status(resolved_reason), + reason=resolved_reason, + phase="in_session", + ) + + +def build_capacity_stop_message() -> Dict[str, Any]: + return build_stop_message( + status=STOP_STATUS_BY_RESOURCE["capacity"], + reason="capacity_exceeded", + resource="capacity", + failed_resources=["capacity"], + phase="pre_ready", + ) + + +def build_bridge_failed_stop_message() -> Dict[str, Any]: + return build_stop_message( + status="stop_bridge_failed", + reason="bridge_failed", + resource="bridge", + phase="in_session", + ) + + +def _build_resource_stop_message(failed_resources: list[str]) -> Dict[str, Any]: + normalized = [name for name in STOP_RESOURCE_PRECEDENCE if name in failed_resources and name != "capacity"] + if not normalized: + raise ValueError("failed_resources must contain at least one non-capacity resource") + primary = normalized[0] + return build_stop_message( + status=STOP_STATUS_BY_RESOURCE[primary], + reason="resource_unhealthy", + resource=primary, + failed_resources=normalized, + phase="pre_ready", + ) + + +def _env_int(name: str, default: int) -> int: + raw = (os.getenv(name, "") or "").strip() + if not raw: + return default + try: + return int(raw) + except ValueError: + return default + + +def _env_float(name: str, default: float) -> float: + raw = (os.getenv(name, "") or "").strip() + if not raw: + return default + try: + return float(raw) + except ValueError: + return default + + +def _env_bool(name: str, default: bool = False) -> bool: + raw = (os.getenv(name, "") or "").strip().lower() + if not raw: + return default + return raw in {"1", "true", "yes", "on"} + + +def _env_bool_default_true(*names: str) -> bool: + for name in names: + raw = (os.getenv(name, "") or "").strip().lower() + if not raw: + continue + if raw in {"0", "false", "no", "off"}: + return False + if raw in {"1", "true", "yes", "on"}: + return True + return True + + +def _strip_text(value: Any) -> str: + return str(value or "").strip() + + +def _env_first(*names: str) -> str: + for name in names: + value = _strip_text(os.getenv(name)) + if value: + return value + return "" + + +def _agent_env_suffix(agent_name: str) -> str: + return "".join( + char if char.isalnum() else "_" + for char in _strip_text(agent_name).upper() + ).strip("_") + + +def _health_desc_value(value: Any) -> str: + if value is None: + return "" + if isinstance(value, str): + return value.strip() + if isinstance(value, (int, float, bool)): + return str(value) + try: + return json.dumps(value, ensure_ascii=False, separators=(",", ":")) + except TypeError: + return str(value).strip() + + +def _health_response_desc(response: httpx.Response) -> str: + try: + payload = response.json() + except ValueError: + payload = None + + if isinstance(payload, Mapping): + for key in ( + "http_cod_desc", + "hhtp_cod_desc", + "cod_desc", + "description", + "message", + "detail", + "error", + "reason", + ): + desc = _health_desc_value(payload.get(key)) + if desc: + return desc + elif payload is None: + text = _strip_text(response.text) + if text: + return text + + return _strip_text(getattr(response, "reason_phrase", "")) or ( + "OK" if 200 <= response.status_code < 300 else "FAIL" + ) + + +def _with_path(url: str, *, path: str) -> str: + parsed = urlsplit(_strip_text(url)) + if not parsed.scheme or not parsed.netloc: + return "" + return urlunsplit((parsed.scheme, parsed.netloc, path, "", "")).rstrip("/") + + +def _derive_http_health_url(url: str) -> str: + return _with_path(url, path="/health") + + +def _derive_websocket_health_url(url: str) -> str: + parsed = urlsplit(_strip_text(url)) + if not parsed.scheme or not parsed.netloc: + return "" + + scheme = parsed.scheme.lower() + if scheme == "ws": + http_scheme = "http" + elif scheme == "wss": + http_scheme = "https" + else: + return "" + + return urlunsplit((http_scheme, parsed.netloc, "/health", "", "")).rstrip("/") + + +def _remote_agent_tls_verify(agent_name: str) -> bool: + suffix = _agent_env_suffix(agent_name) + names: list[str] = [] + if suffix: + names.extend( + [ + f"REMOTE_AGENT_SSE_TLS_VERIFY_{suffix}", + f"REMOTE_AGENT_TLS_VERIFY_{suffix}", + f"REMOTE_AGENT_SSE_VERIFY_TLS_{suffix}", + f"REMOTE_AGENT_VERIFY_TLS_{suffix}", + ] + ) + names.extend( + [ + "REMOTE_AGENT_SSE_TLS_VERIFY", + "REMOTE_AGENT_TLS_VERIFY", + "REMOTE_AGENT_SSE_VERIFY_TLS", + "REMOTE_AGENT_VERIFY_TLS", + ] + ) + return _env_bool_default_true(*names) + + +@dataclass(slots=True) +class ResourceCheck: + status: str + message: str = "" + latency_ms: Optional[int] = None + url: str = "" + http_status_code: Optional[int] = None + http_status_desc: str = "" + + def clone(self) -> "ResourceCheck": + return ResourceCheck( + status=self.status, + message=self.message, + latency_ms=self.latency_ms, + url=self.url, + http_status_code=self.http_status_code, + http_status_desc=self.http_status_desc, + ) + + def is_healthy(self) -> bool: + return self.status in HEALTHY_CHECK_STATUSES + + def as_dict(self) -> Dict[str, Any]: + payload: Dict[str, Any] = {"status": self.status} + if self.message: + payload["message"] = self.message + if self.latency_ms is not None: + payload["latency_ms"] = self.latency_ms + if self.url: + payload["url"] = self.url + return payload + + +@dataclass(slots=True) +class ReadinessReport: + checks: Dict[str, ResourceCheck] + failed_resources: list[str] + active_connections: int + max_connections: int + cached: bool = False + + @property + def status(self) -> str: + return "ok" if not self.failed_resources else "fail" + + def is_healthy(self) -> bool: + return not self.failed_resources + + def clone(self, *, active_connections: Optional[int] = None, max_connections: Optional[int] = None, cached: Optional[bool] = None) -> "ReadinessReport": + return ReadinessReport( + checks={name: check.clone() for name, check in self.checks.items()}, + failed_resources=list(self.failed_resources), + active_connections=self.active_connections if active_connections is None else active_connections, + max_connections=self.max_connections if max_connections is None else max_connections, + cached=self.cached if cached is None else cached, + ) + + def as_dict(self) -> Dict[str, Any]: + return { + "status": self.status, + "cached": self.cached, + "active_connections": self.active_connections, + "max_connections": self.max_connections, + "failed_resources": list(self.failed_resources), + "checks": {name: check.as_dict() for name, check in self.checks.items()}, + } + + def build_stop_message(self) -> Dict[str, Any]: + return _build_resource_stop_message(self.failed_resources) + + +@dataclass(slots=True) +class ServicesHealthReport: + checks: Dict[str, Any] + healthy: bool + + @property + def status(self) -> str: + return "ok" if self.healthy else "fail" + + def is_healthy(self) -> bool: + return self.healthy + + def as_dict(self) -> Dict[str, Any]: + return { + "status": self.status, + "checks": self.checks, + } + + +def _service_health_payload(check: ResourceCheck) -> Dict[str, Any]: + fallback_status = 200 if check.is_healthy() else 500 + fallback_desc = "SUCCESS" if check.is_healthy() else "FAIL" + return { + "htp_cod_status": ( + check.http_status_code + if check.http_status_code is not None + else fallback_status + ), + "hhtp_cod_desc": check.http_status_desc or check.message or fallback_desc, + } + + +class CapacityLease: + def __init__(self, manager: "BridgeReadinessManager") -> None: + self._manager = manager + self._released = False + + async def release(self) -> None: + if self._released: + return + self._released = True + await self._manager._release_capacity() + + +@dataclass(slots=True) +class CapacityReservation: + allowed: bool + lease: Optional[CapacityLease] + active_connections: int + max_connections: int + + +class BridgeReadinessManager: + def __init__(self) -> None: + self._capacity_lock = asyncio.Lock() + self._active_connections = 0 + self._cache_lock = asyncio.Lock() + self._cache: Dict[str, tuple[float, ReadinessReport]] = {} + + def _max_connections(self) -> int: + return max(0, _env_int("TIA_WS_MAX_CONNECTIONS", 0)) + + def _ttl_s(self) -> float: + return max(0.0, _env_float("TIA_RESOURCE_HEALTH_TTL_S", 5.0)) + + def _timeout_s(self) -> float: + return max(0.1, _env_float("TIA_RESOURCE_HEALTH_TIMEOUT_S", 3.0)) + + async def reserve_capacity(self) -> CapacityReservation: + async with self._capacity_lock: + max_connections = self._max_connections() + if max_connections > 0 and self._active_connections >= max_connections: + return CapacityReservation( + allowed=False, + lease=None, + active_connections=self._active_connections, + max_connections=max_connections, + ) + + self._active_connections += 1 + return CapacityReservation( + allowed=True, + lease=CapacityLease(self), + active_connections=self._active_connections, + max_connections=max_connections, + ) + + async def _release_capacity(self) -> None: + async with self._capacity_lock: + if self._active_connections > 0: + self._active_connections -= 1 + + async def capacity_snapshot(self) -> tuple[int, int]: + async with self._capacity_lock: + return self._active_connections, self._max_connections() + + def _cache_key(self, *, agent_name: str, call_config: Mapping[str, Any] | None) -> str: + normalized = normalize_call_config(call_config) + backend_name = resolve_agent_backend_name(normalized, os.getenv("AGENT_BACKEND", "remote_ws")) + stt_provider = (resolve_stt_overrides(normalized).get("provider") or os.getenv("STT_PROVIDER", "internal_http")).strip().lower() + tts_provider = (resolve_tts_overrides(normalized).get("provider") or os.getenv("TTS_PROVIDER", "elevenlabs")).strip().lower() + return "|".join( + [ + _strip_text(agent_name).lower() or "conta", + _strip_text(backend_name).lower() or "remote_ws", + stt_provider or "internal_http", + tts_provider or "elevenlabs", + ] + ) + + async def _get_cached_report(self, key: str) -> Optional[ReadinessReport]: + ttl_s = self._ttl_s() + if ttl_s <= 0: + return None + + async with self._cache_lock: + cached = self._cache.get(key) + if cached is None: + return None + + stored_at, report = cached + if (time.monotonic() - stored_at) > ttl_s: + self._cache.pop(key, None) + return None + + active_connections, max_connections = await self.capacity_snapshot() + return report.clone( + active_connections=active_connections, + max_connections=max_connections, + cached=True, + ) + + async def _store_cached_report(self, key: str, report: ReadinessReport) -> None: + async with self._cache_lock: + self._cache[key] = ( + time.monotonic(), + report.clone(cached=False), + ) + + async def evaluate_resources( + self, + *, + agent_name: str, + call_config: Mapping[str, Any] | None, + ) -> ReadinessReport: + key = self._cache_key(agent_name=agent_name, call_config=call_config) + cached = await self._get_cached_report(key) + if cached is not None: + return cached + + normalized = normalize_call_config(call_config) + agent_runtime, agent_backend, stt, tts = await asyncio.gather( + self.probe_agent_runtime(), + self.probe_agent_backend( + agent_name=agent_name, + call_config=normalized, + ), + self.probe_stt(normalized), + self.probe_tts(normalized), + ) + checks = { + "agent_runtime": agent_runtime, + "agent_backend": agent_backend, + "stt": stt, + "tts": tts, + } + + failed_resources = [ + resource_name + for resource_name in STOP_RESOURCE_PRECEDENCE + if resource_name in checks and not checks[resource_name].is_healthy() + ] + + active_connections, max_connections = await self.capacity_snapshot() + report = ReadinessReport( + checks=checks, + failed_resources=failed_resources, + active_connections=active_connections, + max_connections=max_connections, + cached=False, + ) + await self._store_cached_report(key, report) + return report + + async def evaluate_default_resources(self) -> ReadinessReport: + return await self.evaluate_resources(agent_name="conta", call_config={}) + + async def evaluate_services(self) -> ServicesHealthReport: + ( + agent_runtime, + contas_backend, + oferta_backend, + sofya_stt, + xai_tts, + ) = await asyncio.gather( + self.probe_agent_runtime(), + self.probe_configured_agent_backend_service(agent_name="conta"), + self.probe_configured_agent_backend_service(agent_name="oferta"), + self.probe_configured_stt_service(), + self.probe_configured_xai_tts_service(), + ) + service_checks = [agent_runtime, contas_backend, oferta_backend, sofya_stt, xai_tts] + return ServicesHealthReport( + healthy=all(check.is_healthy() for check in service_checks), + checks={ + "agent_runtime": _service_health_payload(agent_runtime), + "agent_backend": { + "contas": _service_health_payload(contas_backend), + "oferta": _service_health_payload(oferta_backend), + }, + "stt": { + "sofya": _service_health_payload(sofya_stt), + }, + "tts": { + "xAI": _service_health_payload(xai_tts), + }, + }, + ) + + async def probe_configured_agent_backend_service(self, *, agent_name: str) -> ResourceCheck: + normalized_agent = { + "contas": "conta", + "ofertas": "oferta", + }.get(_strip_text(agent_name).lower(), _strip_text(agent_name).lower()) + + if normalized_agent == "conta": + health_url = _env_first( + "REMOTE_AGENT_HEALTH_URL_CONTA", + "REMOTE_AGENT_HEALTH_URL_CONTAS", + "REMOTE_AGENT_HEALTH_URL", + ) + if not health_url: + health_url = _derive_http_health_url( + _env_first("REMOTE_AGENT_SSE_URL_CONTA", "REMOTE_AGENT_SSE_URL_CONTAS") + ) + if not health_url: + health_url = _derive_websocket_health_url( + _env_first("REMOTE_AGENT_WS_URL_CONTA", "REMOTE_AGENT_WS_URL_CONTAS", "REMOTE_AGENT_WS_URL") + ) + elif normalized_agent == "oferta": + health_url = _env_first( + "REMOTE_AGENT_HEALTH_URL_OFERTA", + "REMOTE_AGENT_HEALTH_URL_OFERTAS", + ) + if not health_url: + health_url = _derive_http_health_url( + _env_first("REMOTE_AGENT_SSE_URL_OFERTA", "REMOTE_AGENT_SSE_URL_OFERTAS") + ) + if not health_url: + health_url = _derive_websocket_health_url( + _env_first("REMOTE_AGENT_WS_URL_OFERTA", "REMOTE_AGENT_WS_URL_OFERTAS") + ) + else: + return ResourceCheck(status="fail", message=f"unsupported service agent backend: {agent_name}") + + if not health_url: + return ResourceCheck( + status="fail", + message=f"{normalized_agent} backend health URL is not configured", + ) + + return await self._probe_http_health( + "agent_backend", + health_url, + verify=_remote_agent_tls_verify(normalized_agent), + ) + + async def probe_configured_stt_service(self) -> ResourceCheck: + health_url = _env_first("STT_HEALTH_URL") + if not health_url: + health_url = _derive_http_health_url(_env_first("STT_URL")) + + if not health_url: + return ResourceCheck(status="fail", message="STT health URL is not configured") + + return await self._probe_http_health("stt", health_url) + + async def probe_configured_xai_tts_service(self) -> ResourceCheck: + return await self.probe_tts({"tts": {"provider": "xai"}}) + + async def probe_agent_runtime(self) -> ResourceCheck: + url = f"http://127.0.0.1:{_env_int('AGENT_SERVER_PORT', 18081)}/" + return await self._probe_http_health("agent_runtime", url) + + async def probe_agent_backend( + self, + *, + agent_name: str, + call_config: Mapping[str, Any] | None, + ) -> ResourceCheck: + raw_agent = _strip_text(agent_name).lower() + normalized_agent = { + "contas": "conta", + "ofert": "oferta", + "ofertas": "oferta", + "cobra": "cobranca", + "cobrancas": "cobranca", + "cobrança": "cobranca", + "cobranças": "cobranca", + }.get(raw_agent, raw_agent or "conta") + + backend_name = resolve_agent_backend_name(call_config, os.getenv("AGENT_BACKEND", "remote_ws")).strip().lower() + if backend_name in {"remote_ws_fake", "fake_remote_ws", "ws_fake"}: + return ResourceCheck(status="ok", message="fake backend skips external health probe") + + if backend_name not in {"", "remote_ws", "ws", "websocket", "remote_sse", "sse"}: + return ResourceCheck(status="fail", message=f"unsupported agent backend: {backend_name or 'unknown'}") + + agent_health_envs = { + "conta": ("REMOTE_AGENT_HEALTH_URL_CONTA", "REMOTE_AGENT_HEALTH_URL_CONTAS"), + "oferta": ("REMOTE_AGENT_HEALTH_URL_OFERTA", "REMOTE_AGENT_HEALTH_URL_OFERTAS"), + "cobranca": ( + "REMOTE_AGENT_HEALTH_URL_COBRANCA", + "REMOTE_AGENT_HEALTH_URL_COBRANCAS", + "REMOTE_AGENT_HEALTH_URL_COBRA", + ), + }.get(normalized_agent, ()) + health_url = "" + for env_name in agent_health_envs: + health_url = _strip_text(os.getenv(env_name)) + if health_url: + break + + if normalized_agent == "oferta" and backend_name in {"remote_sse", "sse"} and not health_url: + return ResourceCheck( + status="skipped", + message="oferta backend health probe disabled until health URL is provided", + ) + + if not health_url: + health_url = _strip_text(os.getenv("REMOTE_AGENT_HEALTH_URL")) + + if not health_url: + if backend_name in {"remote_sse", "sse"}: + agent_sse_envs = { + "conta": ("REMOTE_AGENT_SSE_URL_CONTA", "REMOTE_AGENT_SSE_URL_CONTAS"), + "cobranca": ( + "REMOTE_AGENT_SSE_URL_COBRANCA", + "REMOTE_AGENT_SSE_URL_COBRANCAS", + "REMOTE_AGENT_SSE_URL_COBRA", + ), + }.get(normalized_agent, ()) + sse_url = "" + for env_name in agent_sse_envs: + sse_url = _strip_text(os.getenv(env_name)) + if sse_url: + break + health_url = _derive_http_health_url(sse_url) + else: + agent_ws_envs = { + "conta": ("REMOTE_AGENT_WS_URL_CONTA", "REMOTE_AGENT_WS_URL_CONTAS"), + "oferta": ("REMOTE_AGENT_WS_URL_OFERTA", "REMOTE_AGENT_WS_URL_OFERTAS"), + "cobranca": ( + "REMOTE_AGENT_WS_URL_COBRANCA", + "REMOTE_AGENT_WS_URL_COBRANCAS", + "REMOTE_AGENT_WS_URL_COBRA", + ), + }.get(normalized_agent, ()) + ws_url = "" + for env_name in (*agent_ws_envs, "REMOTE_AGENT_WS_URL"): + ws_url = _strip_text(os.getenv(env_name)) + if ws_url: + break + health_url = _derive_websocket_health_url(ws_url) + + if not health_url: + return ResourceCheck(status="fail", message="remote agent health URL is not configured") + + return await self._probe_http_health("agent_backend", health_url) + + async def probe_stt(self, call_config: Mapping[str, Any] | None) -> ResourceCheck: + if _env_bool("TIA_SKIP_STT_READINESS", False): + return ResourceCheck(status="skipped", message="STT readiness disabled by env") + + overrides = resolve_stt_overrides(call_config) + provider = (overrides.get("provider") or os.getenv("STT_PROVIDER", "internal_http")).strip().lower() + + if provider == "fake": + FakeSTT() + return ResourceCheck(status="ok", message="fake STT skips external health probe") + + if provider not in {"", "internal_http"}: + return ResourceCheck(status="fail", message=f"unsupported STT provider: {provider}") + + health_url = _strip_text(os.getenv("STT_HEALTH_URL")) + if not health_url: + health_url = _derive_http_health_url(os.getenv("STT_URL", "")) + + if not health_url: + return ResourceCheck(status="fail", message="STT health URL is not configured") + + return await self._probe_http_health("stt", health_url) + + async def probe_tts(self, call_config: Mapping[str, Any] | None) -> ResourceCheck: + overrides = resolve_tts_overrides(call_config) + provider = (overrides.get("provider") or os.getenv("TTS_PROVIDER", "elevenlabs")).strip().lower() + timeout_s = self._timeout_s() + started_at = time.monotonic() + + try: + if provider == "fake": + tts = FakeProviderTTS() + audio = tts.synthesize_pcm16k(PROBE_TEXT) + self._ensure_audio_bytes(audio) + return ResourceCheck(status="ok", latency_ms=round((time.monotonic() - started_at) * 1000)) + + if provider in {"", "elevenlabs"}: + from livekit.plugins import elevenlabs + + async with aiohttp.ClientSession() as http_session: + tts = elevenlabs.TTS( + model=overrides.get("model_id") or os.getenv("ELEVENLABS_MODEL_ID", ""), + voice_id=overrides.get("voice_id") or os.getenv("ELEVENLABS_VOICE_ID", ""), + api_key=os.getenv("ELEVENLABS_API_KEY", ""), + voice_settings=elevenlabs.VoiceSettings( + stability=0.35, + speed=1.02, + similarity_boost=0.75, + use_speaker_boost=True, + style=0.7, + ), + http_session=http_session, + language="pt", + ) + try: + audio = await asyncio.wait_for(self._collect_livekit_tts_bytes(tts), timeout=timeout_s) + finally: + await tts.aclose() + self._ensure_audio_bytes(audio) + return ResourceCheck(status="ok", latency_ms=round((time.monotonic() - started_at) * 1000)) + + if provider == "xai": + auth_method = _strip_text(os.getenv("XAI_TTS_AUTH_METHOD", "API_KEY")).upper() or "API_KEY" + xai_auth_options: dict[str, str] = {"auth_method": auth_method} + if auth_method == "API_KEY": + api_key = _strip_text(os.getenv("XAI_API_KEY")) + if not api_key: + return ResourceCheck(status="fail", message="missing xAI TTS config: XAI_API_KEY") + xai_auth_options["api_key"] = api_key + readiness_mode = ( + _strip_text(os.getenv(XAI_TTS_READINESS_MODE_ENV, "synthesize")).lower() + or "synthesize" + ) + if readiness_mode not in XAI_TTS_READINESS_MODES: + return ResourceCheck( + status="fail", + message=f"unsupported {XAI_TTS_READINESS_MODE_ENV}: {readiness_mode}", + ) + + async with aiohttp.ClientSession() as http_session: + tts = OraclexAITTS( + voice=overrides.get("voice_id") or os.getenv("XAI_TTS_VOICE", ""), + language=overrides.get("language") or os.getenv("XAI_TTS_LANGUAGE", "pt-BR"), + http_session=http_session, + **xai_auth_options, + ) + try: + if readiness_mode == "connect": + await tts.connect(timeout_s) + return ResourceCheck( + status="ok", + message="xAI TTS websocket connected", + latency_ms=round((time.monotonic() - started_at) * 1000), + ) + audio = await asyncio.wait_for(self._collect_livekit_tts_bytes(tts), timeout=timeout_s) + finally: + await tts.aclose() + self._ensure_audio_bytes(audio) + return ResourceCheck(status="ok", latency_ms=round((time.monotonic() - started_at) * 1000)) + + if provider == "azure": + azure_tts_config, missing_keys = resolve_azure_speech_tts_config(overrides) + if missing_keys: + return ResourceCheck( + status="fail", + message=f"missing Azure TTS config: {', '.join(missing_keys)}", + ) + + azure_impl = (os.getenv("AZURE_TTS_IMPLEMENTATION", "plugin") or "plugin").strip().lower() + if azure_impl == "rest": + tts = AzureRESTTTS( + sample_rate=AZURE_SPEECH_TTS_SAMPLE_RATE, + **azure_tts_config, + ) + try: + audio = await asyncio.wait_for(asyncio.to_thread(tts.synthesize_pcm, PROBE_TEXT), timeout=timeout_s) + finally: + await tts.aclose() + elif azure_impl == "plugin": + from livekit.plugins import azure as livekit_azure + + async with aiohttp.ClientSession() as http_session: + tts = livekit_azure.TTS( + sample_rate=AZURE_SPEECH_TTS_SAMPLE_RATE, + http_session=http_session, + **azure_tts_config, + ) + try: + audio = await asyncio.wait_for(self._collect_livekit_tts_bytes(tts), timeout=timeout_s) + finally: + await tts.aclose() + else: + return ResourceCheck(status="fail", message=f"unsupported AZURE_TTS_IMPLEMENTATION: {azure_impl}") + + self._ensure_audio_bytes(audio) + return ResourceCheck(status="ok", latency_ms=round((time.monotonic() - started_at) * 1000)) + + return ResourceCheck(status="fail", message=f"unsupported TTS provider: {provider}") + except Exception as exc: + return ResourceCheck( + status="fail", + message=_format_probe_error(exc, resource=f"{provider or 'default'} TTS", timeout_s=timeout_s), + latency_ms=round((time.monotonic() - started_at) * 1000), + ) + + async def _collect_livekit_tts_bytes(self, tts: Any) -> bytes: + frame = await tts.synthesize(PROBE_TEXT).collect() + return bytes(getattr(frame, "data", b"")) + + def _ensure_audio_bytes(self, audio: bytes) -> None: + if not audio: + raise RuntimeError("TTS returned empty audio") + + async def _probe_http_health(self, check_name: str, url: str, *, verify: bool = True) -> ResourceCheck: + started_at = time.monotonic() + timeout_s = self._timeout_s() + try: + async with httpx.AsyncClient(follow_redirects=True, verify=verify) as client: + response = await client.get(url, timeout=timeout_s) + latency_ms = round((time.monotonic() - started_at) * 1000) + http_status_desc = _health_response_desc(response) + if 200 <= response.status_code < 300: + return ResourceCheck( + status="ok", + latency_ms=latency_ms, + url=url, + http_status_code=response.status_code, + http_status_desc=http_status_desc, + ) + return ResourceCheck( + status="fail", + message=( + f"{check_name} returned status {response.status_code}: {http_status_desc}" + if http_status_desc + else f"{check_name} returned status {response.status_code}" + ), + latency_ms=latency_ms, + url=url, + http_status_code=response.status_code, + http_status_desc=http_status_desc, + ) + except Exception as exc: + return ResourceCheck( + status="fail", + message=_format_probe_error(exc, resource=check_name, timeout_s=timeout_s), + latency_ms=round((time.monotonic() - started_at) * 1000), + url=url, + ) diff --git a/src/app/ws_gateway/session_audio.py b/src/app/ws_gateway/session_audio.py new file mode 100644 index 0000000..2c78874 --- /dev/null +++ b/src/app/ws_gateway/session_audio.py @@ -0,0 +1,195 @@ +from __future__ import annotations + +import asyncio +import json +import logging +import time +from typing import Any, Awaitable, Callable + +from app.utils.call_timeline import CallTimeline +from app.ws_gateway.readiness import build_completed_stop_message +from app.ws_gateway.session_lifecycle import RoomLifecycleState, close_websocket + + +async def enable_client_audio( + *, + logger: logging.Logger, + timeline: CallTimeline, + client_audio_enabled: asyncio.Event, + room_name: str, + protocol: str, +) -> None: + if client_audio_enabled.is_set(): + return + + client_audio_enabled.set() + logger.info( + "CLIENT_AUDIO_ENABLED | room=%s | protocol=%s", + room_name, + protocol, + ) + timeline.emit("client_audio_enabled") + + +async def notify_client_audio_enabled( + *, + room: Any, + logger: logging.Logger, + timeline: CallTimeline, + room_name: str, + protocol: str, + track_published: asyncio.Event, + agent_ready: asyncio.Event, + lifecycle: RoomLifecycleState, + repeat: bool = False, +) -> None: + await track_published.wait() + await agent_ready.wait() + + last_generation = -1 + while True: + await lifecycle.agent_connected.wait() + generation = lifecycle.agent_generation + if repeat and generation == last_generation: + await asyncio.sleep(0.05) + continue + + if lifecycle.agent_participant is None: + if not repeat: + return + await asyncio.sleep(0.05) + continue + + agent_ident = str(getattr(lifecycle.agent_participant, "identity", "") or "") + if not agent_ident: + if not repeat: + return + await asyncio.sleep(0.05) + continue + + payload = { + "type": "client_audio_enabled", + "room": room_name, + "protocol": protocol, + } + try: + await room.local_participant.publish_data( + json.dumps(payload, ensure_ascii=False), + reliable=True, + destination_identities=[agent_ident], + topic="bridge.control", + ) + logger.info("BRIDGE_CONTROL | sent client_audio_enabled | agent_identity=%s", agent_ident) + timeline.emit( + "bridge_control_sent", + control_type="client_audio_enabled", + agent_identity=agent_ident, + ) + last_generation = generation + except Exception: + logger.debug("BRIDGE_CONTROL send failed", exc_info=True) + + if not repeat: + return + + while lifecycle.agent_generation == generation and lifecycle.agent_connected.is_set(): + await asyncio.sleep(0.05) + + +async def stream_agent_audio_when_ready( + *, + agent_ready: asyncio.Event, + lifecycle: RoomLifecycleState, + agent_q: asyncio.Queue[bytes], + activity: "AudioActivity", + timeline: CallTimeline | None, + stream_agent_audio: Callable[..., Awaitable[None]], + repeat: bool = False, +) -> None: + await agent_ready.wait() + + while True: + await lifecycle.agent_connected.wait() + participant = lifecycle.agent_participant + generation = lifecycle.agent_generation + if participant is None: + if not repeat: + return + await asyncio.sleep(0.05) + continue + + await stream_agent_audio( + agent_q, + participant, + activity, + timeline=timeline, + ) + + if not repeat: + return + + while ( + lifecycle.agent_generation == generation + and lifecycle.agent_participant is participant + and lifecycle.agent_connected.is_set() + ): + await asyncio.sleep(0.05) + + +async def watch_mock_stop_after_first_audio( + *, + ws: Any, + logger: logging.Logger, + timeline: CallTimeline | None, + room_name: str, + protocol: str, + first_agent_audio_sent: asyncio.Event, + activity: "AudioActivity", + agent_q: asyncio.Queue[bytes], + silence_s: float, + reason: str, +) -> None: + await first_agent_audio_sent.wait() + + silence_s = max(0.0, float(silence_s)) + reason = str(reason or "").strip() or "DONE" + logger.info( + "MOCK_STOP_AFTER_FIRST_AUDIO_ARMED | room=%s | protocol=%s | silence_s=%.3f | reason=%s", + room_name, + protocol, + silence_s, + reason, + ) + if timeline is not None: + timeline.emit( + "mock_stop_after_first_audio_armed", + silence_ms=round(silence_s * 1000), + reason=reason, + ) + + poll_interval_s = min(0.1, max(0.02, silence_s / 4 if silence_s > 0 else 0.02)) + while True: + idle_for = time.monotonic() - activity.last_agent_out + if agent_q.empty() and idle_for >= silence_s: + break + await asyncio.sleep(poll_interval_s) + + stop_message = build_completed_stop_message(reason) + try: + await ws.send_text(json.dumps(stop_message, ensure_ascii=False)) + logger.info( + "MOCK_STOP_AFTER_FIRST_AUDIO_SENT | room=%s | protocol=%s | reason=%s", + room_name, + protocol, + reason, + ) + if timeline is not None: + timeline.emit("mock_stop_after_first_audio_sent", reason=reason) + except Exception: + logger.debug("MOCK_STOP_AFTER_FIRST_AUDIO send failed", exc_info=True) + return + + try: + await close_websocket(ws, code=1000, reason=f"mock_stop:{reason}") + except Exception: + logger.debug("MOCK_STOP_AFTER_FIRST_AUDIO close failed", exc_info=True) diff --git a/src/app/ws_gateway/session_bootstrap.py b/src/app/ws_gateway/session_bootstrap.py new file mode 100644 index 0000000..439b449 --- /dev/null +++ b/src/app/ws_gateway/session_bootstrap.py @@ -0,0 +1,140 @@ +from __future__ import annotations + +import logging +import time +import uuid +from dataclasses import dataclass +from typing import Any, Callable, Dict, Mapping + +from app.services.session_context import extract_protocol +from app.utils.call_timeline import CallTimeline +from app.utils.logging import structured_context_from_start_data +from app.ws_gateway.call_config import build_call_config +from app.ws_gateway.session_start import ( + StartSessionContext, + build_remote_agent_context, + should_agent_start_conversation, +) + + +@dataclass(frozen=True, slots=True) +class BridgeSessionBootstrap: + call_id_ged: str + room_name: str + identity: str + token: str + protocol: str + phone_number: str + call_config: Dict[str, Any] + remote_agent_context: Dict[str, Any] + dispatch_metadata: Dict[str, Any] + timeline: CallTimeline + + +def _pick_value(payload: Mapping[str, Any] | None, *keys: str) -> str: + if not isinstance(payload, Mapping): + return "" + + for key in keys: + value = payload.get(key) + if value not in (None, ""): + return str(value).strip() + + lowered = {str(key).lower(): value for key, value in payload.items()} + for key in keys: + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return str(value).strip() + + return "" + + +def build_bridge_session_bootstrap( + *, + start_ctx: StartSessionContext, + app_env: str, + livekit_room: str, + default_protocol: str, + token_factory: Callable[[str, str], str], + logger: logging.Logger, + now_fn: Callable[[], float] = time.time, + uuid_factory: Callable[[], uuid.UUID] = uuid.uuid4, + timeline_origin_unix_ms: int | None = None, +) -> BridgeSessionBootstrap: + call_id_ged = str(start_ctx.data.get("callIdGed") or "")[:64] + call_config = build_call_config(start_ctx.call_config) + remote_agent_context = build_remote_agent_context( + data=start_ctx.data, + agent_data=start_ctx.agent_data, + ) + + base_room = (livekit_room or f"{app_env}-room").strip() + room_name = f"{base_room}-{uuid_factory().hex[:8]}" + identity = f"ws-bridge-{app_env}-{uuid_factory().hex[:8]}" + token = token_factory(identity, room_name) + + protocol = ( + extract_protocol(start_ctx.payload, start_ctx.session_data) + or (default_protocol or "").strip() + or f"WS-{call_id_ged or 'call'}-{int(now_fn())}" + ) + session_id = ( + _pick_value(start_ctx.data, "session_id", "sessionId") + or _pick_value(start_ctx.agent_data, "session_id", "sessionId") + or room_name + ) + phone_number = str( + start_ctx.session_data.get("gsm") + or start_ctx.session_data.get("msisdn") + or start_ctx.session_data.get("phone") + or remote_agent_context.get("GSM") + or "" + ).strip() + structured_context = structured_context_from_start_data(start_ctx.data) + + timeline = CallTimeline( + logger=logger, + component="bridge", + timeline_id=room_name, + protocol=protocol, + room=room_name, + session_id=session_id, + phone_number=phone_number, + origin_unix_ms=timeline_origin_unix_ms, + ) + + dispatch_metadata = { + "env": app_env, + "from": "ws-bridge", + "call_id_ged": call_id_ged, + "protocol": protocol, + "session_id": session_id, + "elegibility": True, + "session_data": start_ctx.session_data, + "intro": start_ctx.intro, + "nudge": start_ctx.nudge, + "agent_starts_conversation": bool( + start_ctx.agent_starts_conversation + or should_agent_start_conversation(start_ctx.data.get("agent")) + ), + "bridge_identity": identity, + "call_config": call_config, + "remote_agent": remote_agent_context, + "timeline_id": room_name, + "timeline_origin_ms": timeline.origin_unix_ms, + "structured_log_context": structured_context.as_metadata(), + "debug_events_enabled": start_ctx.debug_events_enabled, + } + + return BridgeSessionBootstrap( + call_id_ged=call_id_ged, + room_name=room_name, + identity=identity, + token=token, + protocol=protocol, + phone_number=phone_number, + call_config=call_config, + remote_agent_context=remote_agent_context, + dispatch_metadata=dispatch_metadata, + timeline=timeline, + ) diff --git a/src/app/ws_gateway/session_lifecycle.py b/src/app/ws_gateway/session_lifecycle.py new file mode 100644 index 0000000..5917e81 --- /dev/null +++ b/src/app/ws_gateway/session_lifecycle.py @@ -0,0 +1,210 @@ +from __future__ import annotations + +import asyncio +import json +import logging +import time +from dataclasses import dataclass, field +from typing import Any, Awaitable, Callable, Dict, Optional + +from app.utils.call_timeline import CallTimeline +from app.ws_gateway.readiness import build_completed_stop_message, build_stop_message + + +@dataclass(slots=True) +class RoomLifecycleState: + agent_participant: Optional[Any] = None + done_payload: Dict[str, Any] = field(default_factory=dict) + agent_connected: asyncio.Event = field(default_factory=asyncio.Event) + agent_generation: int = 0 + debug_forward_tasks: set[asyncio.Task[Any]] = field(default_factory=set) + + def __post_init__(self) -> None: + if self.agent_participant is not None: + self.agent_connected.set() + if self.agent_generation <= 0: + self.agent_generation = 1 + + +def _close_reason(reason: str, *, limit_bytes: int = 120) -> str: + normalized = " ".join(str(reason or "").strip().split()) + if not normalized: + return "" + + encoded = normalized.encode("utf-8") + if len(encoded) <= limit_bytes: + return normalized + return encoded[:limit_bytes].decode("utf-8", "ignore") + + +async def close_websocket(ws: Any, *, code: int = 1000, reason: str = "") -> None: + close_reason = _close_reason(reason) + try: + await ws.close(code=code, reason=close_reason) + except TypeError: + await ws.close() + + +def register_room_lifecycle_handlers( + *, + room: Any, + logger: logging.Logger, + timeline: CallTimeline, + room_name: str, + protocol: str, + dispatch_started_at: float, + agent_ready: asyncio.Event, + call_done: asyncio.Event, + state: RoomLifecycleState, + is_agent: Callable[[Any], bool], + on_agent_disconnect: Callable[[str, int], Awaitable[None]] | None = None, + ws: Any | None = None, + debug_events_enabled: bool = False, + monotonic_fn: Callable[[], float] = time.monotonic, +) -> None: + @room.on("data_received") + def _on_data_received(data_packet: Any) -> None: + try: + topic = getattr(data_packet, "topic", "") + if topic == "agent.debug": + if not debug_events_enabled or ws is None: + return + raw = data_packet.data + if isinstance(raw, (bytes, bytearray)): + raw = raw.decode("utf-8", "ignore") + message = json.loads(raw) + message["stress_test"] = True + + async def _forward_debug_event() -> None: + try: + await ws.send_text(json.dumps(message, ensure_ascii=False)) + except Exception: + logger.debug("debug event forwarding failed", exc_info=True) + + task = asyncio.create_task( + _forward_debug_event(), name="forward_agent_debug_event" + ) + state.debug_forward_tasks.add(task) + task.add_done_callback(state.debug_forward_tasks.discard) + return + if topic != "agent.stage": + return + + raw = data_packet.data + if isinstance(raw, (bytes, bytearray)): + raw = raw.decode("utf-8", "ignore") + + message = json.loads(raw) + if message.get("stage") == "DONE": + state.done_payload = message + timeline.emit( + "done_packet_received", + reason=message.get("reason", ""), + payload=message, + ) + call_done.set() + except Exception: + logger.exception("DATA_PACKET_FAIL | topic=%s", getattr(data_packet, "topic", "")) + + @room.on("participant_connected") + def _on_participant_connected(participant: Any) -> None: + try: + identity = str(getattr(participant, "identity", "") or "") + if is_agent(participant): + state.agent_participant = participant + state.agent_generation += 1 + state.agent_connected.set() + agent_ready.set() + logger.info( + "LK_AGENT_JOIN | room=%s | agent_identity=%s | dt_s=%.2f", + room_name, + identity or "-", + monotonic_fn() - dispatch_started_at, + ) + timeline.emit( + "agent_join", + agent_identity=identity or "", + dispatch_dt_ms=round((monotonic_fn() - dispatch_started_at) * 1000), + ) + else: + logger.debug("LK_PARTICIPANT_JOIN | room=%s | identity=%s", room_name, identity or "-") + except Exception: + logger.debug("participant_connected handler failed", exc_info=True) + + @room.on("participant_disconnected") + def _on_participant_disconnected(participant: Any) -> None: + try: + identity = str(getattr(participant, "identity", "") or "") + if is_agent(participant): + current_identity = str(getattr(state.agent_participant, "identity", "") or "") + cleared_current_agent = False + if state.agent_participant is participant or current_identity == identity: + state.agent_participant = None + state.agent_generation += 1 + state.agent_connected.clear() + cleared_current_agent = True + logger.info("LK_AGENT_LEAVE | room=%s | agent_identity=%s", room_name, identity or "-") + timeline.emit("agent_leave", agent_identity=identity or "") + if cleared_current_agent and not call_done.is_set() and on_agent_disconnect is not None: + generation = state.agent_generation + + async def _notify_disconnect() -> None: + try: + await on_agent_disconnect(identity, generation) + except Exception: + logger.exception( + "AGENT_DISCONNECT_HANDLER_FAIL | room=%s | protocol=%s | agent_identity=%s", + room_name, + protocol, + identity or "-", + ) + + asyncio.create_task(_notify_disconnect(), name="agent_disconnect_handler") + else: + logger.debug("LK_PARTICIPANT_LEAVE | room=%s | identity=%s", room_name, identity or "-") + except Exception: + logger.debug("participant_disconnected handler failed", exc_info=True) + + +async def watch_call_done( + *, + call_done: asyncio.Event, + state: RoomLifecycleState, + ws: Any, + logger: logging.Logger, + room_name: str, + protocol: str, + timeline: CallTimeline, +) -> None: + await call_done.wait() + + payload = state.done_payload or {} + reason = str(payload.get("reason") or "DONE") + logger.info("CALL_DONE | room=%s | protocol=%s | reason=%s", room_name, protocol, reason) + timeline.emit("call_done", reason=reason) + + if state.debug_forward_tasks: + await asyncio.gather(*tuple(state.debug_forward_tasks), return_exceptions=True) + + stop_status = str(payload.get("status") or "").strip() + if stop_status: + stop_message = build_stop_message( + status=stop_status, + reason=reason, + resource=str(payload.get("resource") or "").strip(), + failed_resources=list(payload.get("failed_resources") or []), + phase=str(payload.get("phase") or "in_session").strip(), + ) + else: + stop_message = build_completed_stop_message(reason) + + try: + await ws.send_text(json.dumps(stop_message, ensure_ascii=False)) + timeline.emit("stop_sent", reason=reason) + except Exception: + pass + + try: + await close_websocket(ws, code=1000, reason=f"call_done:{reason}") + except Exception: + pass diff --git a/src/app/ws_gateway/session_start.py b/src/app/ws_gateway/session_start.py new file mode 100644 index 0000000..e087cd0 --- /dev/null +++ b/src/app/ws_gateway/session_start.py @@ -0,0 +1,326 @@ +from __future__ import annotations + +import json +import logging +from dataclasses import dataclass +from typing import Any, Dict, Mapping + +from fastapi import WebSocket + +from app.common.call_config import is_fake_agent_stress_test, resolve_fake_agent_overrides + + +@dataclass(frozen=True, slots=True) +class StartSessionContext: + payload: Dict[str, Any] + data: Dict[str, Any] + agent_data: Dict[str, Any] + audio_format: Dict[str, Any] + call_config: Dict[str, Any] + session_data: Dict[str, Any] + intro: str + nudge: str + agent_starts_conversation: bool = False + stress_test: bool = False + debug_events_enabled: bool = False + + +def _pick_value(payload: Mapping[str, Any], *keys: str) -> str: + for key in keys: + value = payload.get(key) + if value not in (None, ""): + return str(value).strip() + + lowered = {str(key).lower(): value for key, value in payload.items()} + for key in keys: + value = lowered.get(str(key).lower()) + if value not in (None, ""): + return str(value).strip() + + return "" + + +def _normalize_agent_name(value: Any) -> str: + raw = str(value or "").strip().lower() + aliases = { + "conta": "conta", + "contas": "conta", + "ofert": "oferta", + "oferta": "oferta", + "ofertas": "oferta", + "cobra": "cobranca", + "cobranca": "cobranca", + "cobrancas": "cobranca", + } + return aliases.get(raw, raw) + + +def should_agent_start_conversation(agent_name: str) -> bool: + # O backend remoto é responsável pela saudação inicial em todos os agentes. + del agent_name + return True + + +def _mapping_or_error(payload: Mapping[str, Any], key: str) -> Dict[str, Any]: + value = payload.get(key) + if not isinstance(value, Mapping): + raise RuntimeError(f"Campo obrigatório inválido: {key}") + return dict(value) + + +def _optional_mapping(payload: Mapping[str, Any], key: str) -> Dict[str, Any]: + value = payload.get(key) + if value in (None, ""): + return {} + if not isinstance(value, Mapping): + raise RuntimeError(f"Campo opcional inválido: {key}") + return dict(value) + + +def _required_value(payload: Mapping[str, Any], *keys: str) -> str: + value = _pick_value(payload, *keys) + if not value: + joined = ", ".join(keys) + raise RuntimeError(f"Campo obrigatório ausente: {joined}") + return value + + +def parse_start_payload(payload: Mapping[str, Any]) -> StartSessionContext: + if payload.get("type") != "start": + raise RuntimeError( + f"Primeira mensagem esperada type='start', veio: {payload.get('type')}" + ) + + start_payload = dict(payload) + data = _mapping_or_error(start_payload, "data") + agent_data = _optional_mapping(data, "agentData") + audio_format = _optional_mapping(start_payload, "audioFormat") + call_config = _optional_mapping(start_payload, "callConfig") + try: + resolve_fake_agent_overrides(call_config) + stress_test = is_fake_agent_stress_test(call_config) + except ValueError as exc: + raise RuntimeError(f"callConfig invalido: {exc}") from exc + # Structured diagnostics may contain high-volume, implementation-specific + # data. Forward them to the WS client only for an explicitly scripted fake + # agent call; client-provided debug flags alone must not enable forwarding. + debug_events_enabled = stress_test + + data["agent"] = _normalize_agent_name(_required_value(data, "agent")) + _required_value(data, "ani") + _required_value(data, "gsm") + _required_value(data, "routerCallKeyDay") + _required_value(data, "routerCallKey") + _required_value(data, "callIdGed") + session_id = _required_value(data, "session_id", "sessionId") + data["session_id"] = session_id + + if data["agent"] == "conta": + _required_value( + agent_data, + "idFatura", + "current_invoice_number", + "currentInvoiceNumber", + ) + elif data["agent"] == "oferta": + data["protocolo"] = _required_value( + data, + "protocolo", + "protocolNumber", + "protocol", + ) + + session_data = dict(data) + if agent_data: + session_data["agentData"] = dict(agent_data) + if audio_format: + session_data["audioFormat"] = dict(audio_format) + if session_data.get("gsm"): + session_data["msisdn"] = str(session_data["gsm"]).strip() + session_data["phone"] = str(session_data["gsm"]).strip() + + start_payload["data"] = dict(data) + agent_starts_conversation = should_agent_start_conversation(data["agent"]) + intro = "" + nudge = "Alô, você ainda está aí?" + + return StartSessionContext( + payload=start_payload, + data=data, + agent_data=agent_data, + audio_format=audio_format, + call_config=call_config, + session_data=session_data, + intro=intro, + nudge=nudge, + agent_starts_conversation=agent_starts_conversation, + stress_test=stress_test, + debug_events_enabled=debug_events_enabled, + ) + + +def parse_transferencia_session_id_payload(payload: Mapping[str, Any]) -> str: + if payload.get("type") != "transferencia_session_id": + raise RuntimeError( + "Mensagem esperada type='transferencia_session_id', " + f"veio: {payload.get('type')}" + ) + + data = _mapping_or_error(payload, "data") + return _required_value(data, "session_id", "sessionId") + + +def _json_for_log(payload: Any) -> str: + try: + return json.dumps(payload, ensure_ascii=False, indent=2, default=str) + except TypeError: + return str(payload) + + +def _redact_fake_responses(payload: Mapping[str, Any]) -> Dict[str, Any]: + sanitized = dict(payload) + call_config = sanitized.get("callConfig") + if not isinstance(call_config, Mapping): + return sanitized + safe_call_config = dict(call_config) + agent_fake = safe_call_config.get("agentFake") + if isinstance(agent_fake, Mapping): + safe_agent_fake = dict(agent_fake) + raw = safe_agent_fake.pop("responses", None) + if raw not in (None, ""): + lengths = [len(item.strip()) for item in str(raw).split(";")] + safe_agent_fake["responseCount"] = len(lengths) + safe_agent_fake["responseLengths"] = lengths + safe_call_config["agentFake"] = safe_agent_fake + sanitized["callConfig"] = safe_call_config + return sanitized + + +def _log_start_request(logger: logging.Logger | None, payload: Mapping[str, Any]) -> None: + if logger is None: + return + + logger.info( + "WS_START_REQUEST\npayload=%s\ncall_config=%s", + _json_for_log(_redact_fake_responses(payload)), + _json_for_log(_redact_fake_responses(payload).get("callConfig") or {}), + ) + + +async def recv_start_message( + ws: WebSocket, + logger: logging.Logger | None = None, +) -> StartSessionContext: + pending_session_id = "" + + while True: + raw_text = await ws.receive_text() + payload = json.loads(raw_text) + + if payload.get("type") == "transferencia_session_id": + pending_session_id = parse_transferencia_session_id_payload(payload) + continue + + if payload.get("type") == "start" and pending_session_id: + start_payload = dict(payload) + raw_data = start_payload.get("data") + if isinstance(raw_data, Mapping): + data = dict(raw_data) + if not _pick_value(data, "session_id", "sessionId"): + data["session_id"] = pending_session_id + start_payload["data"] = data + payload = start_payload + + _log_start_request(logger, payload) + return parse_start_payload(payload) + + + +def build_remote_agent_context( + *, + data: Mapping[str, Any], + agent_data: Mapping[str, Any], +) -> Dict[str, Any]: + agent_name = _normalize_agent_name(_pick_value(data, "agent")) + + context = { + "agent": agent_name, + "RouterCallKeyDay": _pick_value(data, "routerCallKeyDay"), + "RouterCallKey": _pick_value(data, "routerCallKey"), + "ANI": _pick_value(data, "ani"), + "GSM": _pick_value(data, "gsm"), + "msisdn": _pick_value(data, "gsm"), + "callIdGed": _pick_value(data, "callIdGed"), + } + + explicit_session_id = ( + _pick_value(data, "session_id", "sessionId") + or _pick_value(agent_data, "session_id", "sessionId") + ) + if explicit_session_id: + context["session_id"] = explicit_session_id + context["sessionId"] = explicit_session_id + + explicit_message_id = ( + _pick_value(data, "message_id", "messageId") + or _pick_value(agent_data, "message_id", "messageId") + ) + if explicit_message_id: + context["message_id"] = explicit_message_id + + explicit_protocol_id = ( + _pick_value(data, "protocol_id", "protocolId") + or _pick_value(agent_data, "protocol_id", "protocolId") + ) + protocol_number = ( + explicit_protocol_id + or _pick_value(data, "protocolo", "protocolNumber", "protocol") + or _pick_value(agent_data, "protocolo", "protocolNumber", "protocol") + ) + if protocol_number: + context["protocolo"] = protocol_number + context["protocolNumber"] = protocol_number + if explicit_protocol_id: + context["protocol_id"] = explicit_protocol_id + + channel_id = ( + _pick_value(data, "channelId", "channel_id") + or _pick_value(agent_data, "channelId", "channel_id") + ) + if channel_id: + context["channelId"] = channel_id + + asset_id = ( + _pick_value(data, "assetId", "asset_id") + or _pick_value(agent_data, "assetId", "asset_id") + ) + if asset_id: + context["assetId"] = asset_id + + if agent_name == "conta": + invoice_number = ( + _pick_value( + agent_data, + "current_invoice_number", + "currentInvoiceNumber", + "idFatura", + ) + or _pick_value( + data, + "current_invoice_number", + "currentInvoiceNumber", + "idFatura", + ) + ) + if invoice_number: + context["ID_FATURA"] = invoice_number + context["current_invoice_number"] = invoice_number + + channel = ( + _pick_value(agent_data, "channel", "Channel", "canal") + or _pick_value(data, "channel", "Channel", "canal") + ) + if channel: + context["channel"] = channel + + return context diff --git a/src/app/ws_gateway/voice_client.html b/src/app/ws_gateway/voice_client.html new file mode 100644 index 0000000..55dfaee --- /dev/null +++ b/src/app/ws_gateway/voice_client.html @@ -0,0 +1,1750 @@ + + + + + + TIM TIA Voice Client + + + +
+
+
+
+ +
+
TIM | TIA Voice Console
+
Imagine as possibilidades
+
+
+

Client

+

+ Console de homologação para capturar o microfone, enviar PCM16 ao bridge em + /ws/agent, acompanhar os eventos da sessão e reproduzir o áudio de resposta no + navegador com uma apresentação alinhada à marca TIM. +

+
+
+ +
+
+
+
+

Configuração da Chamada

+

+ Dados de sessão, roteamento e providers usados na homologação da jornada de voz. +

+
+
+
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ +

+ Os campos abaixo mudam conforme o agente remoto selecionado. +

+
+
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ +
+ Observacoes +

+ O contrato de `start` usa `type`, `data`, `audioFormat` e `callConfig`. Para homologacao local sem + agent remoto de pe, use `remote_ws_fake` com `fake` em STT e TTS. Quando o agent de contas estiver + publicado via SSE, selecione `remote_sse`. Os campos adicionais aparecem conforme o agent remoto. +

+
+
+ +
+
+
+

Controle e Telemetria

+

+ Disparo da chamada, status em tempo real e eventos trocados com o bridge. +

+
+
+
+ + + +
+
+
idle
+
mic off
+
audio off
+
+
+
+ Sua voz + +
microfone em espera
+
+
+ Resposta + +
saida em espera
+
+
+ +

Log da Sessão

+
+
+
+
+ + + + diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/adapters/__init__.py b/tests/adapters/__init__.py new file mode 100644 index 0000000..9d48db4 --- /dev/null +++ b/tests/adapters/__init__.py @@ -0,0 +1 @@ +from __future__ import annotations diff --git a/tests/adapters/test_audio_gain.py b/tests/adapters/test_audio_gain.py new file mode 100644 index 0000000..27256f9 --- /dev/null +++ b/tests/adapters/test_audio_gain.py @@ -0,0 +1,134 @@ +from __future__ import annotations + +import struct + +import numpy as np + +from app.livekit.adapters.audio_gain import ( + GainEmitter, + SoftClipGain, + tts_output_gain_from_env, +) + + +def _pcm(*samples: int) -> bytes: + return struct.pack("<" + "h" * len(samples), *samples) + + +def _samples(pcm: bytes) -> list[int]: + return list(struct.unpack("<" + "h" * (len(pcm) // 2), pcm)) + + +def test_gain_1_0_is_disabled_and_passthrough() -> None: + g = SoftClipGain(gain=1.0) + assert g.enabled is False + pcm = _pcm(1000, -2000, 3000) + assert g.process(pcm) == pcm + + +def test_normal_level_is_boosted_near_linear() -> None: + g = SoftClipGain(gain=2.0) # ceiling default -1 dBFS + # sinal baixo (~ -30 dBFS): boost deve ser praticamente 2x + out = _samples(g.process(_pcm(1000, -1000))) + assert abs(out[0] - 2000) <= 40 + assert abs(out[1] + 2000) <= 40 + + +def test_hot_peaks_never_clip_past_ceiling() -> None: + ceiling = 0.891 # ~ -1 dBFS + g = SoftClipGain(gain=2.0, ceiling=ceiling) + limit = int(ceiling * 32768) + 1 + # picos quentes que, com 2x linear, estourariam o fundo de escala + out = _samples(g.process(_pcm(30000, -30000, 25000, -25000))) + assert all(abs(s) <= limit for s in out), out + # e continua monotonicamente crescente (sem wraparound/inversao de fase) + assert out[0] > 0 and out[1] < 0 + + +def test_monotonic_transfer_curve() -> None: + g = SoftClipGain(gain=2.0) + xs = list(range(0, 32000, 1000)) + ys = [_samples(g.process(_pcm(x)))[0] for x in xs] + assert all(b >= a for a, b in zip(ys, ys[1:])), ys + + +def test_odd_length_bytes_do_not_crash() -> None: + g = SoftClipGain(gain=2.0) + pcm = _pcm(1000, -1000) + b"\x7f" # 1 byte solto + out = g.process(pcm) + assert len(out) == len(pcm) + assert out[-1:] == b"\x7f" + + +def test_empty_input() -> None: + assert SoftClipGain(gain=2.0).process(b"") == b"" + + +def test_env_loader_defaults_to_disabled(monkeypatch) -> None: + monkeypatch.delenv("TTS_OUTPUT_GAIN", raising=False) + monkeypatch.delenv("TTS_OUTPUT_CEILING_DBFS", raising=False) + g = tts_output_gain_from_env() + assert g.gain == 1.0 + assert g.enabled is False + + +def test_env_loader_reads_gain_and_ceiling(monkeypatch) -> None: + monkeypatch.setenv("TTS_OUTPUT_GAIN", "2.0") + monkeypatch.setenv("TTS_OUTPUT_CEILING_DBFS", "-6") + g = tts_output_gain_from_env() + assert g.gain == 2.0 + assert abs(g.ceiling - 10 ** (-6 / 20.0)) < 1e-6 + + +class _FakeEmitter: + def __init__(self) -> None: + self.pushed: list[bytes] = [] + self.initialized = False + self.flushed = False + + def initialize(self, **kwargs) -> None: + self.initialized = True + + def push(self, data: bytes) -> None: + self.pushed.append(data) + + def flush(self) -> None: + self.flushed = True + + +def test_gain_emitter_transforms_push_and_forwards_rest() -> None: + inner = _FakeEmitter() + em = GainEmitter(inner, SoftClipGain(gain=2.0)) + + em.initialize(sample_rate=24000) + em.push(_pcm(1000, -1000)) + em.flush() + + assert inner.initialized is True + assert inner.flushed is True + assert len(inner.pushed) == 1 + # o que chegou ao emitter real foi amplificado + out = _samples(inner.pushed[0]) + assert abs(out[0] - 2000) <= 40 + + +def test_gain_emitter_matches_numpy_reference() -> None: + gain = SoftClipGain(gain=2.0, ceiling=0.891) + pcm = _pcm(500, -12000, 28000, -31000, 0) + ref_x = np.frombuffer(pcm, dtype=" None: + monkeypatch.setenv("TTS_OUTPUT_GAIN", "nan") + monkeypatch.setenv("TTS_OUTPUT_CEILING_DBFS", "inf") + gain = tts_output_gain_from_env() + assert gain.gain == 1.0 + + monkeypatch.setenv("TTS_OUTPUT_GAIN", "-2") + monkeypatch.setenv("TTS_OUTPUT_CEILING_DBFS", "-200") + gain = tts_output_gain_from_env() + assert gain.ceiling == 0.05 + assert gain.gain == 0.0 diff --git a/tests/adapters/test_azure_rest_tts.py b/tests/adapters/test_azure_rest_tts.py new file mode 100644 index 0000000..a97e94a --- /dev/null +++ b/tests/adapters/test_azure_rest_tts.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +import unittest +from unittest import mock + +import httpx + +from app.livekit.adapters.azure_rest_tts import AzureRESTTTS + + +class AzureRESTTTSTests(unittest.TestCase): + def test_request_endpoints_keep_explicit_tts_path(self) -> None: + tts = AzureRESTTTS( + voice="pt-BR-FranciscaNeural", + speech_key="key-123", + speech_region="brazilsouth", + speech_endpoint="https://speech.example.cognitiveservices.azure.com/tts/cognitiveservices/v1", + ) + + self.assertEqual( + tts._request_endpoints(), + ["https://speech.example.cognitiveservices.azure.com/tts/cognitiveservices/v1"], + ) + + def test_synthesize_pcm_tries_voice_path_after_base_404_for_custom_voice(self) -> None: + tts = AzureRESTTTS( + voice="pt-BR-FranciscaNeural", + speech_key="key-123", + speech_endpoint="https://speech.example.cognitiveservices.azure.com/cognitiveservices/v1", + deployment_id="deployment-42", + ) + + base_url = "https://speech.example.cognitiveservices.azure.com/cognitiveservices/v1?deploymentId=deployment-42" + voice_url = "https://speech.example.cognitiveservices.azure.com/voice/cognitiveservices/v1?deploymentId=deployment-42" + + response_404 = httpx.Response(404, request=httpx.Request("POST", base_url)) + response_ok = httpx.Response(200, content=b"\x00\x00", request=httpx.Request("POST", voice_url)) + + with mock.patch.object(httpx.Client, "post", side_effect=[response_404, response_ok]) as mocked_post: + audio = tts.synthesize_pcm("teste") + + self.assertEqual(audio, b"\x00\x00") + self.assertEqual([call.args[0] for call in mocked_post.call_args_list], [base_url, voice_url]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/adapters/test_fake_tts.py b/tests/adapters/test_fake_tts.py new file mode 100644 index 0000000..7c3cb1e --- /dev/null +++ b/tests/adapters/test_fake_tts.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +import importlib +import sys +import types +import unittest +from types import SimpleNamespace + + +def _install_fake_livekit() -> None: + if "livekit.agents" in sys.modules: + return + + try: + importlib.import_module("livekit.agents") + return + except ImportError: + pass + + try: + livekit_pkg = importlib.import_module("livekit") + except ImportError: + livekit_pkg = types.ModuleType("livekit") + livekit_pkg.__path__ = [] + sys.modules["livekit"] = livekit_pkg + + agents_module = types.ModuleType("livekit.agents") + types_module = types.ModuleType("livekit.agents.types") + + class TTSCapabilities: + def __init__(self, *, streaming: bool, aligned_transcript: bool) -> None: + self.streaming = streaming + self.aligned_transcript = aligned_transcript + + class AudioEmitter: + def __init__(self) -> None: + self._data = bytearray() + self.sample_rate = 0 + self.num_channels = 0 + + def initialize(self, *, request_id: str, sample_rate: int, num_channels: int, mime_type: str) -> None: + self.request_id = request_id + self.sample_rate = sample_rate + self.num_channels = num_channels + self.mime_type = mime_type + + def push(self, data: bytes) -> None: + if data: + self._data.extend(data) + + def flush(self) -> None: + return None + + def snapshot(self): + return SimpleNamespace( + sample_rate=self.sample_rate, + num_channels=self.num_channels, + data=bytes(self._data), + ) + + class BaseTTS: + def __init__(self, *, capabilities: TTSCapabilities, sample_rate: int, num_channels: int) -> None: + self.capabilities = capabilities + self.sample_rate = sample_rate + self.num_channels = num_channels + + class BaseChunkedStream: + def __init__(self, *, tts: BaseTTS, input_text: str, conn_options) -> None: + self._tts = tts + self._input_text = input_text + self._conn_options = conn_options + + async def collect(self): + emitter = AudioEmitter() + await self._run(emitter) + return emitter.snapshot() + + class APIConnectOptions: + def __init__(self, **kwargs) -> None: + self.kwargs = kwargs + + tts_module = SimpleNamespace( + TTS=BaseTTS, + TTSCapabilities=TTSCapabilities, + ChunkedStream=BaseChunkedStream, + AudioEmitter=AudioEmitter, + ) + + agents_module.tts = tts_module + agents_module.utils = SimpleNamespace(shortuuid=lambda: "req-test") + + types_module.APIConnectOptions = APIConnectOptions + types_module.DEFAULT_API_CONNECT_OPTIONS = APIConnectOptions() + + setattr(livekit_pkg, "agents", agents_module) + sys.modules["livekit.agents"] = agents_module + sys.modules["livekit.agents.types"] = types_module + + +_install_fake_livekit() + +from app.livekit.adapters.fake_tts import FakeTTS + + +class FakeLiveKitTTSTests(unittest.IsolatedAsyncioTestCase): + async def test_synthesize_collects_pcm_audio(self) -> None: + tts = FakeTTS() + stream = tts.synthesize("teste fake") + frame = await stream.collect() + + self.assertEqual(frame.sample_rate, 16000) + self.assertEqual(frame.num_channels, 1) + self.assertGreater(len(frame.data), 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/adapters/test_xai_tts.py b/tests/adapters/test_xai_tts.py new file mode 100644 index 0000000..e7ddf4e --- /dev/null +++ b/tests/adapters/test_xai_tts.py @@ -0,0 +1,407 @@ +from __future__ import annotations + +import asyncio +import base64 +import json +import os +import unittest +from types import SimpleNamespace +from unittest import mock + +import aiohttp + +from app.livekit.adapters import xai_tts as xai_tts_module +from app.livekit.adapters.xai_tts import ( + AUTH_METHOD_API_KEY, + DEFAULT_LANGUAGE, + DEFAULT_VOICE, + OraclexAITTS, +) + + +def _event(event_type: str, **values: object) -> SimpleNamespace: + return SimpleNamespace( + type=aiohttp.WSMsgType.TEXT, + data=json.dumps({"type": event_type, **values}), + ) + + +class _FakeWebSocket: + def __init__(self, events: list[SimpleNamespace]) -> None: + self.events = list(events) + self.sent: list[dict[str, object]] = [] + self.closed = False + + def exception(self): + return None + + async def send_str(self, payload: str) -> None: + self.sent.append(json.loads(payload)) + + async def receive(self) -> SimpleNamespace: + if not self.events: + await asyncio.sleep(60) + return self.events.pop(0) + + async def close(self) -> None: + self.closed = True + + +class _Emitter: + def __init__(self) -> None: + self.audio = bytearray() + + def push(self, payload: bytes) -> None: + self.audio.extend(payload) + + +class _Stream: + _segment_id = "segment-test" + + def _mark_started(self) -> None: + return None + + def _note_provider_ttfb(self, _provider_ttfb: float) -> None: + return None + + +class _EventOwner: + def __init__(self) -> None: + self.events: list[tuple[str, dict[str, object]]] = [] + + def _take_initial_greeting_capture(self, _text: str) -> None: + return None + + def emit(self, event_name: str, event: dict[str, object]) -> None: + self.events.append((event_name, event)) + + +class XAITTSUpgradeTests(unittest.IsolatedAsyncioTestCase): + def test_underflow_limit_defaults_to_one_second(self) -> None: + with mock.patch.dict(os.environ, {}, clear=True): + self.assertEqual(xai_tts_module._underflow_error_ms(), 1000) + self.assertEqual(xai_tts_module._turn_total_timeout_s(), 60.0) + + def test_estimated_pcm_balance_uses_first_pcm_release_without_prebuffer(self) -> None: + self.assertEqual( + xai_tts_module._estimated_pcm_balance_s( + pcm_duration_s=1.0, first_pcm_released_at=100.0, now=101.125 + ), + -0.125, + ) + + def _connection( + self, events: list[SimpleNamespace], owner: _EventOwner | None = None + ): + options = xai_tts_module._TTSOptions( + base_url="wss://example.test/tts", + voice=DEFAULT_VOICE, + language=DEFAULT_LANGUAGE, + ) + auth = xai_tts_module._AuthOptions( + method=AUTH_METHOD_API_KEY, + api_key="key-123", + ) + connection = xai_tts_module._Connection( + opts=options, + auth=auth, + session=object(), + owner=owner, + ) + connection._ws = _FakeWebSocket(events) + connection._note_activity() + return connection + + async def _synthesize(self, events: list[SimpleNamespace]): + connection = self._connection(events) + emitter = _Emitter() + result = await connection.synthesize_turn( + "nova fala", + output_emitter=emitter, + stream=_Stream(), + timeout=1.0, + turn_index=0, + connection_reused=True, + ) + return connection, emitter, result + + def test_legacy_public_name_and_websocket_url_are_preserved(self) -> None: + tts = OraclexAITTS(api_key="key-123", websocket_url="wss://legacy.test/tts") + + self.assertIs(xai_tts_module.TTS, OraclexAITTS) + self.assertEqual(tts._opts.base_url, "wss://legacy.test/tts") + self.assertEqual(tts._opts.voice, DEFAULT_VOICE) + self.assertEqual(tts._opts.language, DEFAULT_LANGUAGE) + self.assertEqual(tts.model, DEFAULT_VOICE) + + def test_api_key_auth_and_iam_auth_validation_are_available(self) -> None: + with mock.patch.dict(os.environ, {"XAI_API_KEY": "env-key"}, clear=True): + tts = OraclexAITTS() + + self.assertEqual(tts._auth.method, AUTH_METHOD_API_KEY) + self.assertEqual(tts._auth.api_key, "env-key") + self.assertEqual( + xai_tts_module._request_headers(tts._auth, "wss://example.test/tts"), + {"Authorization": "Bearer env-key"}, + ) + + with self.assertRaisesRegex(ValueError, "compartment_id"): + OraclexAITTS(auth_method="INSTANCE_PRINCIPAL") + + def test_cached_greeting_read_failure_falls_back_to_tts(self) -> None: + class BrokenCache: + def __init__(self) -> None: + self.key = SimpleNamespace(digest="cache-key") + self.discarded = [] + + def key_for(self, **_kwargs): + return self.key + + def has(self, _key) -> bool: + return True + + def frames(self, _key): + raise FileNotFoundError("cached WAV disappeared") + + def discard(self, key) -> None: + self.discarded.append(key) + + cache = BrokenCache() + tts = OraclexAITTS( + api_key="key-123", + websocket_url="wss://example.test/tts", + initial_greeting_audio_cache=cache, + initial_greeting_agent="conta", + ) + + self.assertIsNone(tts.initial_greeting_audio("Olá, como posso ajudar?")) + self.assertEqual(cache.discarded, [cache.key]) + self.assertIs(tts._initial_greeting_capture_key, cache.key) + + def test_cached_greeting_hit_is_logged(self) -> None: + cache = mock.Mock() + cache_key = SimpleNamespace(digest="cache-key") + cached_audio = object() + cache.key_for.return_value = cache_key + cache.has.return_value = True + cache.frames.return_value = cached_audio + tts = OraclexAITTS( + api_key="key-123", + websocket_url="wss://example.test/tts", + initial_greeting_audio_cache=cache, + initial_greeting_agent="conta", + ) + + with mock.patch.object(xai_tts_module, "_runtime_logger") as runtime_logger: + self.assertIs( + cached_audio, tts.initial_greeting_audio("Olá, como posso ajudar?") + ) + + runtime_logger.return_value.info.assert_called_once_with( + "INITIAL_GREETING_AUDIO_CACHE_HIT | key=%s", "cache-key" + ) + cache.discard.assert_not_called() + + async def test_connect_exposes_a_deterministic_prewarm_operation(self) -> None: + tts = OraclexAITTS(api_key="key-123") + with mock.patch.object(tts, "_current_connection", new=mock.AsyncMock()) as current_connection: + await tts.connect(1.5) + + current_connection.assert_awaited_once_with(1.5) + await tts.aclose() + + async def test_stream_merges_text_chunks_into_one_provider_turn(self) -> None: + greeting_chunks = ( + "Olá! Eu sou a Especialista em Contas e vou ajudar você a entender a sua fatura. ", + "Posso explicar valores, detalhar serviços e itens eventuais, identificar cobranças que você não reconhece e, se for o caso, realizar ajustes necessários ou solicitações relacionadas à sua conta. ", + "Então vamos lá, me conte o que você gostaria de entender ou resolver na sua conta.", + ) + greeting = "".join(greeting_chunks) + pcm = b"\x01\x00" * 4800 + websocket = _FakeWebSocket( + [ + _event("audio.clear"), + _event("audio.delta", delta=base64.b64encode(pcm).decode()), + _event("audio.done", trace_id="trace-chunked"), + ] + ) + session = SimpleNamespace( + ws_connect=mock.AsyncMock(return_value=websocket), + ) + cache_key = SimpleNamespace(digest="greeting-key", text=greeting) + cache = mock.Mock() + cache.key_for.return_value = cache_key + cache.has.return_value = False + cache.store_pcm = mock.AsyncMock() + + provider = OraclexAITTS( + api_key="key-123", + websocket_url="wss://example.test/tts", + http_session=session, + initial_greeting_audio_cache=cache, + initial_greeting_agent="contas", + ) + + self.assertIsNone(provider.initial_greeting_audio(greeting)) + + stream = provider.stream() + for chunk in greeting_chunks: + stream.push_text(chunk) + stream.end_input() + try: + audio_events = [event async for event in stream] + finally: + await stream.aclose() + await provider.aclose() + + self.assertIsInstance(stream, xai_tts_module.SynthesizeStream) + self.assertEqual(b"".join(event.frame.data.tobytes() for event in audio_events), pcm) + self.assertEqual( + [message["type"] for message in websocket.sent], + ["text.clear", "text.delta", "text.done"], + ) + self.assertEqual(websocket.sent[1]["delta"], greeting) + cache.store_pcm.assert_awaited_once_with(cache_key, pcm) + + def test_text_sanitization_is_preserved(self) -> None: + self.assertEqual( + xai_tts_module._sanitize_tts_text( + "TIM_GAMES_KIDS_MES custa R$ 14,99 no dia 29/05/26; a/b?" + ), + "TIM GAMES KIDS MES custa R 14,99 no dia 29/05/26; a ou b?", + ) + + async def test_turn_requires_clear_ack_before_emitting_audio(self) -> None: + payload = base64.b64encode(b"novo").decode() + connection, emitter, result = await self._synthesize( + [_event("audio.clear"), _event("audio.delta", delta=payload), _event("audio.done", trace_id="trace-1")] + ) + + self.assertEqual(bytes(emitter.audio), b"novo") + self.assertEqual( + [message["type"] for message in connection._ws.sent], + ["text.clear", "text.delta", "text.done"], + ) + self.assertEqual(result.trace_id, "trace-1") + self.assertEqual(result.timing.discarded_messages, []) + + async def test_residual_audio_is_discarded_before_clear_ack(self) -> None: + old_payload = base64.b64encode(b"velho").decode() + new_payload = base64.b64encode(b"novo").decode() + _connection, emitter, result = await self._synthesize( + [ + _event("audio.delta", delta=old_payload), + _event("audio.done"), + _event("audio.clear"), + _event("audio.delta", delta=new_payload), + _event("audio.done", trace_id="trace-new"), + ] + ) + + self.assertEqual(bytes(emitter.audio), b"novo") + self.assertEqual(result.trace_id, "trace-new") + self.assertEqual(result.timing.discarded_messages, ["audio.delta", "audio.done"]) + self.assertEqual(result.timing.clear_discarded_message_count, 2) + self.assertEqual(result.timing.clear_discarded_audio_bytes, len(b"velho")) + + async def test_unexpected_clear_after_boundary_fails_and_retires_socket(self) -> None: + payload = base64.b64encode(b"parcial").decode() + connection = self._connection( + [_event("audio.clear"), _event("audio.delta", delta=payload), _event("audio.clear")] + ) + emitter = _Emitter() + + with self.assertRaisesRegex(xai_tts_module._XAIPartialAudioFailure, "unexpected_audio_clear"): + await connection.synthesize_turn( + "fala", + output_emitter=emitter, + stream=_Stream(), + timeout=1.0, + turn_index=0, + connection_reused=True, + ) + + self.assertEqual(bytes(emitter.audio), b"parcial") + self.assertIsNone(connection._ws) + + async def test_first_frame_timeout_is_configurable(self) -> None: + connection = self._connection([_event("audio.clear")]) + connection._session = SimpleNamespace( + ws_connect=mock.AsyncMock(return_value=_FakeWebSocket([])) + ) + emitter = _Emitter() + with mock.patch.dict(os.environ, {"TTS_FIRST_FRAME_TIMEOUT_S": "0.01"}, clear=False): + with self.assertRaisesRegex(xai_tts_module.APIConnectionError, "timed out before audio"): + await connection.synthesize_turn( + "fala", + output_emitter=emitter, + stream=_Stream(), + timeout=1.0, + turn_index=0, + connection_reused=True, + ) + + self.assertIsNone(connection._ws) + + async def test_audio_done_without_pcm_retries_once_on_same_socket(self) -> None: + pcm = b"\x01\x00" * 480 + connection, emitter, result = await self._synthesize([ + _event("audio.clear"), + _event("audio.done"), + _event("audio.clear"), + _event("audio.delta", delta=base64.b64encode(pcm).decode()), + _event("audio.done", trace_id="trace-after-empty"), + ]) + self.assertEqual(bytes(emitter.audio), pcm) + self.assertEqual(result.timing.attempts, 2) + self.assertFalse(connection._ws.closed) + self.assertEqual( + [item["type"] for item in connection._ws.sent], + ["text.clear", "text.delta", "text.done"] * 2, + ) + + async def test_partial_audio_resync_discards_socket_without_replay(self) -> None: + pcm = b"\x01\x00" * 480 + owner = _EventOwner() + connection = self._connection([ + _event("audio.clear"), + _event("audio.delta", delta=base64.b64encode(pcm).decode()), + ], owner=owner) + emitter = _Emitter() + websocket = connection._ws + with self.assertRaisesRegex(xai_tts_module._XAIPartialAudioFailure, "socket_resynchronized=0"): + await connection.synthesize_turn("fala", output_emitter=emitter, stream=_Stream(), timeout=1.0, turn_index=0, connection_reused=True) + self.assertEqual([item["type"] for item in websocket.sent].count("text.delta"), 1) + self.assertEqual([item["type"] for item in websocket.sent].count("text.clear"), 2) + self.assertEqual(len(owner.events), 1) + event_name, event = owner.events[0] + self.assertEqual(event_name, "xai_tts_turn_failed") + self.assertEqual(event["segment_id"], "segment-test") + self.assertEqual(event["reason"], "underflow_error") + self.assertEqual(event["xai_micro_underflows"], 1) + self.assertGreaterEqual(event["max_playout_underrun_0ms"], 10) + self.assertEqual(event["attempts"], 1) + self.assertGreater(event["pcm_bytes"], 0) + self.assertGreater(event["pcm_duration_ms"], 0) + self.assertTrue(event["connection_reused"]) + self.assertFalse(event["reconnected"]) + self.assertIsNone(connection._ws) + + async def test_continuous_underflow_discards_socket_without_replaying_text(self) -> None: + connection = self._connection([ + _event("audio.clear"), + _event("audio.delta", delta=base64.b64encode(b"\x01\x00").decode()), + ]) + emitter = _Emitter() + websocket = connection._ws + with mock.patch.dict(os.environ, {"TTS_UNDERFLOW_ERROR_MS": "10", "TTS_FIRST_FRAME_TIMEOUT_S": "0.01"}, clear=False): + with self.assertRaisesRegex(xai_tts_module._XAIPartialAudioFailure, "underflow_error") as raised: + await connection.synthesize_turn("fala", output_emitter=emitter, stream=_Stream(), timeout=1.0, turn_index=0, connection_reused=True) + self.assertIn("xai_micro_underflows=1", str(raised.exception)) + self.assertIn("xai_avg_underrun_ms=", str(raised.exception)) + self.assertEqual([item["type"] for item in websocket.sent].count("text.delta"), 1) + self.assertIsNone(connection._ws) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/config/__init__.py b/tests/config/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/config/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/config/test_azure_speech.py b/tests/config/test_azure_speech.py new file mode 100644 index 0000000..e080cc5 --- /dev/null +++ b/tests/config/test_azure_speech.py @@ -0,0 +1,116 @@ +from __future__ import annotations + +import unittest + +from app.livekit.azure_speech import resolve_azure_speech_tts_config + + +class AzureSpeechTTSTests(unittest.TestCase): + def test_resolve_uses_region_voice_and_optional_language(self) -> None: + config, missing = resolve_azure_speech_tts_config( + {}, + environ={ + "AZURE_SPEECH_KEY": "key-123", + "AZURE_SPEECH_REGION": "brazilsouth", + "AZURE_SPEECH_VOICE": "pt-BR-FranciscaNeural", + "AZURE_SPEECH_LANGUAGE": "pt-BR", + }, + ) + + self.assertEqual(missing, []) + self.assertEqual(config["speech_key"], "key-123") + self.assertEqual(config["speech_region"], "brazilsouth") + self.assertIsNone(config["speech_endpoint"]) + self.assertEqual(config["voice"], "pt-BR-FranciscaNeural") + self.assertEqual(config["language"], "pt-BR") + self.assertIsNone(config["deployment_id"]) + + def test_resolve_keeps_explicit_custom_endpoint_for_standard_voice(self) -> None: + config, missing = resolve_azure_speech_tts_config( + {}, + environ={ + "AZURE_SPEECH_KEY": "key-123", + "AZURE_SPEECH_REGION": "brazilsouth", + "AZURE_SPEECH_ENDPOINT": "https://speech.example.cognitiveservices.azure.com/", + "AZURE_SPEECH_VOICE": "pt-BR-FranciscaNeural", + }, + ) + + self.assertEqual(missing, []) + self.assertEqual(config["speech_region"], "brazilsouth") + self.assertEqual( + config["speech_endpoint"], + "https://speech.example.cognitiveservices.azure.com/tts/cognitiveservices/v1", + ) + + def test_resolve_prefers_endpoint_and_maps_legacy_override_names(self) -> None: + config, missing = resolve_azure_speech_tts_config( + { + "voice_id": "pt-BR-FranciscaNeural", + "model_id": "deployment-42", + }, + environ={ + "AZURE_SPEECH_KEY": "key-123", + "AZURE_SPEECH_REGION": "brazilsouth", + "AZURE_SPEECH_ENDPOINT": "https://speech.example.cognitiveservices.azure.com/", + }, + ) + + self.assertEqual(missing, []) + self.assertEqual(config["deployment_id"], "deployment-42") + self.assertEqual( + config["speech_endpoint"], + "https://speech.example.cognitiveservices.azure.com/voice/cognitiveservices/v1", + ) + self.assertEqual(config["speech_region"], "brazilsouth") + + def test_resolve_accepts_host_alias_and_auth_token(self) -> None: + config, missing = resolve_azure_speech_tts_config( + {}, + environ={ + "AZURE_SPEECH_AUTH_TOKEN": "token-123", + "AZURE_SPEECH_HOST": "https://speech.example.cognitiveservices.azure.com/", + "AZURE_SPEECH_VOICE": "pt-BR-FranciscaNeural", + }, + ) + + self.assertEqual(missing, []) + self.assertIsNone(config["speech_key"]) + self.assertEqual(config["speech_auth_token"], "token-123") + self.assertEqual( + config["speech_endpoint"], + "https://speech.example.cognitiveservices.azure.com/tts/cognitiveservices/v1", + ) + + def test_resolve_keeps_full_endpoint_path(self) -> None: + config, missing = resolve_azure_speech_tts_config( + {}, + environ={ + "AZURE_SPEECH_KEY": "key-123", + "AZURE_SPEECH_ENDPOINT": "https://speech.example.cognitiveservices.azure.com/tts/cognitiveservices/v1", + "AZURE_SPEECH_VOICE": "pt-BR-FranciscaNeural", + }, + ) + + self.assertEqual(missing, []) + self.assertEqual( + config["speech_endpoint"], + "https://speech.example.cognitiveservices.azure.com/tts/cognitiveservices/v1", + ) + + def test_resolve_reports_missing_required_settings(self) -> None: + config, missing = resolve_azure_speech_tts_config({}, environ={}) + + self.assertEqual(config, {}) + self.assertEqual( + missing, + [ + "AZURE_SPEECH_VOICE", + "AZURE_SPEECH_ENDPOINT|AZURE_SPEECH_HOST|AZURE_SPEECH_REGION", + "AZURE_SPEECH_KEY|AZURE_SPEECH_AUTH_TOKEN", + ], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/config/test_call_config.py b/tests/config/test_call_config.py new file mode 100644 index 0000000..3208ebb --- /dev/null +++ b/tests/config/test_call_config.py @@ -0,0 +1,254 @@ +from __future__ import annotations + +import unittest + +from app.livekit.call_config import ( + normalize_call_config, + resolve_fake_agent_overrides, + resolve_agent_backend_name, + resolve_stt_overrides, + resolve_tts_overrides, + resolve_vad_logging_overrides, + resolve_vad_overrides, + resolve_ws_overrides, +) +from app.ws_gateway.call_config import build_call_config + + +class CallConfigTests(unittest.TestCase): + def test_build_and_normalize_call_config(self) -> None: + payload = { + "agentBackend": "remote_ws_fake", + "stt": { + "provider": "internal_http", + "language": "pt-BR", + "configOverride": "{\"processor\":{\"strategy\":\"faster_default\"}}", + "minProbSingleWord": "0.15", + "disableVosk": True, + }, + "tts": { + "provider": "elevenlabs", + "voiceId": "voice-123", + "modelId": "model-456", + }, + "vad": { + "minSpeechDuration": 0.04, + "activationThreshold": 0.18, + "deactivationThreshold": 0.10, + "minSilenceDuration": 0.7, + "prefixPaddingDuration": 0.75, + "preBackendWaitNoticeFastOnVadPause": True, + "deferredInterruptionMinAudioMs": 1350, + "deferredInterruptionEnabled": False, + }, + "vadLogging": { + "logDecisions": True, + "logActivity": False, + "activityMinProbability": 0.03, + }, + "ws": { + "outputGain": 1.2, + "audioInputBacklogShedEnabled": True, + "audioInputBacklogShedThresholdMs": 700, + "audioInputBacklogShedKeepMs": 250, + "audioInputLatencyMetricsEnabled": True, + "audioInputLatencyAlertMs": 900, + "audioInputLatencyLogIntervalS": 10.5, + "livekitAudioSourceQueueSizeMs": 400, + "livekitAudioSourceClearOnShed": False, + "audioInputBacklogEnergyShedEnabled": True, + "audioInputBacklogEnergyShedMaxExcessMs": 600, + "audioInputBacklogSilenceDbfs": -60, + }, + "agentFake": { + "delayMs": 2500, + "responses": ( + "Esta e a primeira resposta simulada com tamanho intermediario;" + "Esta e a resposta final simulada encerrando o atendimento" + ), + }, + } + + built = build_call_config(payload) + normalized = normalize_call_config(payload) + + self.assertEqual(built, normalized) + self.assertEqual(resolve_agent_backend_name(payload, "remote_ws"), "remote_ws_fake") + self.assertEqual(resolve_stt_overrides(payload)["language"], "pt-BR") + self.assertEqual(resolve_stt_overrides(payload)["disable_vosk"], "True") + self.assertEqual( + resolve_stt_overrides(payload)["config_override"], + "{\"processor\":{\"strategy\":\"faster_default\"}}", + ) + self.assertEqual(resolve_tts_overrides(payload)["voice_id"], "voice-123") + self.assertEqual(resolve_vad_overrides(payload)["min_speech_duration"], "0.04") + self.assertEqual(resolve_vad_overrides(payload)["activation_threshold"], "0.18") + self.assertEqual( + resolve_vad_overrides(payload)["pre_backend_wait_notice_fast_on_vad_pause"], + "True", + ) + self.assertEqual( + resolve_vad_overrides(payload)["deferred_interruption_min_audio_ms"], + "1350", + ) + self.assertEqual( + resolve_vad_overrides(payload)["deferred_interruption_enabled"], + "False", + ) + self.assertEqual(resolve_vad_logging_overrides(payload)["log_decisions"], "True") + self.assertEqual(resolve_vad_logging_overrides(payload)["log_activity"], "False") + self.assertEqual(resolve_ws_overrides(payload)["output_gain"], "1.2") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_backlog_shed_enabled"], "True") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_backlog_shed_threshold_ms"], "700") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_backlog_shed_keep_ms"], "250") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_latency_metrics_enabled"], "True") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_latency_alert_ms"], "900") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_latency_log_interval_s"], "10.5") + self.assertEqual(resolve_ws_overrides(payload)["livekit_audio_source_queue_size_ms"], "400") + self.assertEqual(resolve_ws_overrides(payload)["livekit_audio_source_clear_on_shed"], "False") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_backlog_energy_shed_enabled"], "True") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_backlog_energy_shed_max_excess_ms"], "600") + self.assertEqual(resolve_ws_overrides(payload)["audio_in_backlog_silence_dbfs"], "-60") + fake = resolve_fake_agent_overrides(payload) + self.assertEqual(fake["delay_ms"], 2500) + self.assertEqual(len(fake["responses"]), 2) + + def test_empty_payload_falls_back_to_defaults(self) -> None: + payload = {} + + self.assertEqual( + build_call_config(payload), + { + "agent_backend": "", + "stt": { + "provider": "", + "language": "", + "api_key": "", + "initial_prompt": "", + "config_override": "", + "min_prob_single_word": "", + "disable_vosk": "", + }, + "tts": { + "provider": "", + "voice_id": "", + "model_id": "", + "language": "", + }, + "vad": { + "min_speech_duration": "", + "activation_threshold": "", + "deactivation_threshold": "", + "min_silence_duration": "", + "prefix_padding_duration": "", + "pre_backend_wait_notice_fast_on_vad_pause": "", + "deferred_interruption_min_audio_ms": "", + "deferred_interruption_enabled": "", + }, + "vad_logging": { + "log_decisions": "", + "log_activity": "", + "activity_min_probability": "", + }, + "ws": { + "output_gain": "", + "audio_in_backlog_shed_enabled": "", + "audio_in_backlog_shed_threshold_ms": "", + "audio_in_backlog_shed_keep_ms": "", + "audio_in_latency_metrics_enabled": "", + "audio_in_latency_alert_ms": "", + "audio_in_latency_log_interval_s": "", + "livekit_audio_source_queue_size_ms": "", + "livekit_audio_source_clear_on_shed": "", + "audio_in_backlog_energy_shed_enabled": "", + "audio_in_backlog_energy_shed_max_excess_ms": "", + "audio_in_backlog_silence_dbfs": "", + }, + "agent_fake": {"delay_ms": "", "responses": ""}, + }, + ) + self.assertEqual(resolve_agent_backend_name(payload, "remote_ws"), "remote_ws") + + def test_fake_responses_are_trimmed_and_default_delay_is_applied(self) -> None: + payload = { + "agentBackend": "remote_ws_fake", + "agentFake": { + "responses": ( + " Primeira resposta simulada com comprimento intermediario ;" + " Segunda resposta simulada encerrando corretamente a chamada " + ) + }, + } + + fake = resolve_fake_agent_overrides(payload) + + self.assertEqual(fake["delay_ms"], 2500) + self.assertEqual( + fake["responses"], + [ + "Primeira resposta simulada com comprimento intermediario", + "Segunda resposta simulada encerrando corretamente a chamada", + ], + ) + + def test_fake_responses_reject_invalid_contract(self) -> None: + valid = "Esta resposta simulada possui tamanho intermediario adequado" + invalid_cases = [ + {"agentBackend": "remote_ws", "agentFake": {"responses": f"{valid};{valid}"}}, + {"agentBackend": "remote_ws_fake", "agentFake": {"responses": valid}}, + {"agentBackend": "remote_ws_fake", "agentFake": {"responses": f"{valid};;{valid}"}}, + {"agentBackend": "remote_ws_fake", "agentFake": {"responses": f"curta;{valid}"}}, + { + "agentBackend": "remote_ws_fake", + "agentFake": {"responses": f"{valid};{valid}", "delayMs": 180001}, + }, + ] + + for payload in invalid_cases: + with self.subTest(payload=payload), self.assertRaises(ValueError): + resolve_fake_agent_overrides(payload) + + def test_disable_vosk_and_fake_backend_are_normalized_per_call(self) -> None: + payload = { + "agentBackend": "remote_ws_fake", + "stt": {"provider": "internal_http", "disableVosk": True}, + } + + self.assertEqual(resolve_agent_backend_name(payload, "remote_ws"), "remote_ws_fake") + self.assertEqual(resolve_stt_overrides(payload)["disable_vosk"], "True") + + def test_fast_vad_pause_override_lives_with_vad_overrides(self) -> None: + payload = {"vad": {"preBackendWaitNoticeFastOnVadPause": True}} + + self.assertEqual( + resolve_vad_overrides(payload)["pre_backend_wait_notice_fast_on_vad_pause"], + "True", + ) + + def test_deferred_interruption_min_audio_override_lives_with_vad_overrides(self) -> None: + payload = {"vad": {"deferredInterruptionMinAudioMs": 1250}} + + self.assertEqual( + resolve_vad_overrides(payload)["deferred_interruption_min_audio_ms"], + "1250", + ) + + def test_deferred_interruption_min_audio_accepts_root_env_style_key(self) -> None: + payload = {"DEFERRED_INTERRUPTION_MIN_AUDIO_MS": 1400} + + self.assertEqual( + resolve_vad_overrides(payload)["deferred_interruption_min_audio_ms"], + "1400", + ) + + def test_deferred_interruption_enabled_accepts_root_env_style_key(self) -> None: + payload = {"DEFERRED_INTERRUPTION_ENABLED": False} + + self.assertEqual( + resolve_vad_overrides(payload)["deferred_interruption_enabled"], + "False", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..bfa9e99 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + + +ROOT_DIR = Path(__file__).resolve().parent.parent +SRC_DIR = ROOT_DIR / "src" + +if str(SRC_DIR) not in sys.path: + sys.path.insert(0, str(SRC_DIR)) diff --git a/tests/livekit/__init__.py b/tests/livekit/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/livekit/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/livekit/test_agent_finalization.py b/tests/livekit/test_agent_finalization.py new file mode 100644 index 0000000..f203ef5 --- /dev/null +++ b/tests/livekit/test_agent_finalization.py @@ -0,0 +1,21 @@ +from __future__ import annotations + +from app.livekit.policies.agent_finalization import ( + final_stop_from_agent_result, + stop_status_for_agent_result_type, +) + + +def test_stop_status_for_agent_result_type_uses_conta_contract() -> None: + assert stop_status_for_agent_result_type("resolvido") == "stop_resolvido_e_finalizado" + assert stop_status_for_agent_result_type("nao_resolvido") == "stop_nao_resolvido" + assert stop_status_for_agent_result_type("resolvido_outros_assuntos") == "stop_outro_assunto" + assert stop_status_for_agent_result_type("outros_assuntos") == "stop_outro_assunto" + assert stop_status_for_agent_result_type("erro_falha_sistema") == "stop_falha_sistema" + assert stop_status_for_agent_result_type("erro_no_match") == "stop_no_match" + + +def test_final_stop_from_agent_result_ignores_non_terminal_result_types() -> None: + assert final_stop_from_agent_result({"type": "final", "content": "texto"}) is None + assert final_stop_from_agent_result({"content": "texto"}) is None + assert final_stop_from_agent_result(None) is None diff --git a/tests/livekit/test_bridge_gateway.py b/tests/livekit/test_bridge_gateway.py new file mode 100644 index 0000000..10528ef --- /dev/null +++ b/tests/livekit/test_bridge_gateway.py @@ -0,0 +1,68 @@ +from __future__ import annotations + +import asyncio +import json + +from app.livekit.adapters.bridge_gateway import BridgeGateway + + +class _Participant: + def __init__(self) -> None: + self.calls = [] + + async def publish_data(self, payload, **kwargs) -> None: + self.calls.append((json.loads(payload), kwargs)) + + +class _Room: + name = "room-load-1" + + def __init__(self) -> None: + self.local_participant = _Participant() + + +def test_publish_debug_event_targets_originating_bridge() -> None: + async def _run() -> None: + room = _Room() + gateway = BridgeGateway( + room=room, + bridge_identity="bridge-load-1", + protocol="LOAD-1", + stress_test=True, + ) + + await gateway.publish_debug_event( + "stt.completed", duration_ms=123, text="texto reconhecido" + ) + + payload, kwargs = room.local_participant.calls[0] + assert payload["type"] == "debug_event" + assert payload["event"] == "stt.completed" + assert payload["stress_test"] is True + assert payload["data"] == { + "duration_ms": 123, + "text": "texto reconhecido", + } + assert kwargs == { + "reliable": True, + "destination_identities": ["bridge-load-1"], + "topic": "agent.debug", + } + + asyncio.run(_run()) + + +def test_publish_debug_event_is_suppressed_outside_stress_test() -> None: + async def _run() -> None: + room = _Room() + gateway = BridgeGateway( + room=room, + bridge_identity="bridge-regular-1", + protocol="REGULAR-1", + ) + + await gateway.publish_debug_event("stt.completed", duration_ms=123) + + assert room.local_participant.calls == [] + + asyncio.run(_run()) diff --git a/tests/livekit/test_compat.py b/tests/livekit/test_compat.py new file mode 100644 index 0000000..b8b82a4 --- /dev/null +++ b/tests/livekit/test_compat.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +import types + +import pytest + +from app.livekit.compat import patch_inference_executor_is_alive + + +def test_patch_inference_executor_is_alive_handles_closed_process(monkeypatch) -> None: + fake_module = types.SimpleNamespace() + + class FakeInferenceProcExecutor: + def is_alive(self) -> bool: + raise ValueError("process object is closed") + + fake_module.InferenceProcExecutor = FakeInferenceProcExecutor + + monkeypatch.setattr( + "app.livekit.compat.import_module", + lambda name: fake_module, + ) + + patched = patch_inference_executor_is_alive() + + assert patched is True + assert FakeInferenceProcExecutor().is_alive() is False + + +def test_patch_inference_executor_is_alive_preserves_other_value_errors(monkeypatch) -> None: + fake_module = types.SimpleNamespace() + + class FakeInferenceProcExecutor: + def is_alive(self) -> bool: + raise ValueError("unexpected failure") + + fake_module.InferenceProcExecutor = FakeInferenceProcExecutor + + monkeypatch.setattr( + "app.livekit.compat.import_module", + lambda name: fake_module, + ) + + patch_inference_executor_is_alive() + + with pytest.raises(ValueError, match="unexpected failure"): + FakeInferenceProcExecutor().is_alive() + + +def test_patch_inference_executor_is_alive_is_idempotent(monkeypatch) -> None: + fake_module = types.SimpleNamespace() + + class FakeInferenceProcExecutor: + calls = 0 + + def is_alive(self) -> bool: + type(self).calls += 1 + raise ValueError("process object is closed") + + fake_module.InferenceProcExecutor = FakeInferenceProcExecutor + + monkeypatch.setattr( + "app.livekit.compat.import_module", + lambda name: fake_module, + ) + + first = patch_inference_executor_is_alive() + second = patch_inference_executor_is_alive() + + assert first is True + assert second is False + assert FakeInferenceProcExecutor().is_alive() is False + assert FakeInferenceProcExecutor.calls == 1 diff --git a/tests/livekit/test_fake_remote_ws_adapter.py b/tests/livekit/test_fake_remote_ws_adapter.py new file mode 100644 index 0000000..f57ff13 --- /dev/null +++ b/tests/livekit/test_fake_remote_ws_adapter.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import unittest + +from app.livekit.adapters.agent_backend import BackendReply +from app.livekit.adapters.fake_remote_ws_adapter import FakeRemoteWSAdapter + + +class FakeRemoteWSAdapterTests(unittest.IsolatedAsyncioTestCase): + RESPONSES = ( + "Primeira resposta deterministica com comprimento intermediario", + "Segunda resposta deterministica preparando a formalizacao", + "Terceira resposta deterministica encerrando o atendimento", + ) + + async def test_run_returns_mocked_reply_with_transcribed_text(self) -> None: + adapter = FakeRemoteWSAdapter( + intro="oi", + request_context={ + "agent": "oferta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0001", + "ANI": "3133334444", + "GSM": "31999999999", + "callIdGed": "GED-123", + }, + ) + await adapter.prepare(True, "PRT-123") + + reply = await adapter.run({"text": "quero detalhes da oferta"}) + + self.assertEqual(reply.stage, "ARGUMENTATION") + self.assertFalse(reply.done) + self.assertIn("quero detalhes da oferta", reply.text.lower()) + + async def test_end_service_once_returns_done_without_endpoint_dependency(self) -> None: + adapter = FakeRemoteWSAdapter( + intro="oi", + request_context={"agent": "conta", "GSM": "5511999999999"}, + ) + await adapter.prepare(True, "PRT-999") + + reply = await adapter.end_service_once() + + self.assertEqual( + reply, + BackendReply( + stage="DONE", + text="Atendimento simulado de conta encerrado. Obrigado.", + done=True, + export_payload={ + "type": "final", + "content": "Atendimento simulado de conta encerrado. Obrigado.", + "tool_calls": [], + "result": [{"status": "ok", "reason": "fake_done"}], + }, + ), + ) + + async def test_scripted_responses_are_sequential_and_end_idempotently(self) -> None: + adapter = FakeRemoteWSAdapter( + intro="saudacao normal", + request_context={"agent": "oferta"}, + delay_ms=0, + responses=self.RESPONSES, + ) + await adapter.prepare(True, "PRT-SEQUENTIAL") + + first = await adapter.run({"text": "texto do STT que nao influencia a resposta"}) + second = await adapter.run({"text": "outro texto arbitrario reconhecido"}) + third = await adapter.run({"text": "ultimo texto arbitrario reconhecido"}) + repeated = await adapter.run({"text": "texto posterior ao encerramento"}) + + self.assertEqual( + [first.stage, second.stage, third.stage], + ["ARGUMENTATION", "FORMALIZATION", "DONE"], + ) + self.assertEqual([first.text, second.text, third.text], list(self.RESPONSES)) + self.assertTrue(third.done) + self.assertIs(repeated, third) + + async def test_scripted_sequence_is_isolated_per_adapter_session(self) -> None: + first_call = FakeRemoteWSAdapter(intro="oi", delay_ms=0, responses=self.RESPONSES) + second_call = FakeRemoteWSAdapter(intro="oi", delay_ms=0, responses=self.RESPONSES) + await first_call.prepare(True, "PRT-1") + await second_call.prepare(True, "PRT-2") + + await first_call.run("fala um") + first_reply_second_call = await second_call.run("fala independente") + + self.assertEqual(first_reply_second_call.text, self.RESPONSES[0]) + self.assertEqual(first_reply_second_call.stage, "ARGUMENTATION") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/livekit/test_initial_greeting_audio_cache.py b/tests/livekit/test_initial_greeting_audio_cache.py new file mode 100644 index 0000000..f2bca9f --- /dev/null +++ b/tests/livekit/test_initial_greeting_audio_cache.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +import tempfile +import unittest +import wave +from pathlib import Path +from unittest import mock + +from app.livekit.runtime.initial_greeting_audio_cache import InitialGreetingAudioCache + + +class InitialGreetingAudioCacheTests(unittest.IsolatedAsyncioTestCase): + def test_cache_defaults_are_enabled_and_bounded(self) -> None: + with mock.patch.dict("os.environ", {}, clear=True): + cache = InitialGreetingAudioCache() + + self.assertTrue(cache.enabled) + self.assertEqual(cache._max_chars, 1_000) + self.assertEqual(cache._ttl_s, 3_600) + self.assertEqual(cache._max_entries, 32) + self.assertEqual(cache._max_bytes, 16 * 1024 * 1024) + self.assertEqual(cache._max_agent_variants, 1) + self.assertEqual(cache._disable_ttl_s, 300) + key = cache.key_for( + agent="conta", + text="Olá, como posso ajudar?", + provider="xAI", + voice="ara", + language="pt-BR", + sample_rate=24000, + ) + self.assertIsNotNone(key) + assert key is not None + self.assertEqual(key.agent, "conta") + + async def test_store_writes_atomic_pcm_wav_for_matching_key(self) -> None: + with tempfile.TemporaryDirectory() as directory: + with mock.patch.dict( + "os.environ", + { + "INITIAL_GREETING_AUDIO_CACHE_ENABLED": "1", + "INITIAL_GREETING_AUDIO_CACHE_DIR": directory, + }, + clear=True, + ): + cache = InitialGreetingAudioCache() + key = cache.key_for( + agent="conta", + text="Olá, como posso ajudar?", + provider="xAI", + voice="ara", + language="pt-BR", + sample_rate=24000, + ) + assert key is not None + await cache.store_pcm(key, b"\x00\x00" * 240) + + path = Path(directory) / f"{key.digest}.wav" + self.assertTrue(path.is_file()) + with wave.open(str(path), "rb") as rendered: + self.assertEqual(rendered.getnchannels(), 1) + self.assertEqual(rendered.getsampwidth(), 2) + self.assertEqual(rendered.getframerate(), 24000) + self.assertEqual(rendered.getnframes(), 240) diff --git a/tests/livekit/test_remote_agent_sse_adapter.py b/tests/livekit/test_remote_agent_sse_adapter.py new file mode 100644 index 0000000..5cf7da3 --- /dev/null +++ b/tests/livekit/test_remote_agent_sse_adapter.py @@ -0,0 +1,1129 @@ +from __future__ import annotations + +import json +import os +import unittest +import uuid +from unittest import mock + +import httpx + +from app.livekit.adapters.agent_backend import BackendReply +from app.livekit.adapters.backend_factory import build_agent_backend +from app.livekit.adapters.remote_agent_sse_adapter import RemoteAgentSSEAdapter + + +def _assert_uuid(testcase: unittest.TestCase, value: object) -> str: + text = str(value or "") + testcase.assertEqual(str(uuid.UUID(text)), text) + return text + + +def _sse_event(event_name: str, data: object) -> list[str]: + return [ + f"event: {event_name}", + f"data: {json.dumps(data)}", + "", + ] + + +class _FakeStreamResponse: + def __init__(self, lines, *, method: str, url: str) -> None: + self._lines = list(lines) + self.request = httpx.Request(method, url) + self.status_code = 200 + self.closed = False + + def raise_for_status(self) -> None: + return None + + async def aiter_lines(self): + for line in self._lines: + yield line + + async def aclose(self) -> None: + self.closed = True + + +class _FakeStreamContextManager: + def __init__(self, response: _FakeStreamResponse) -> None: + self._response = response + + async def __aenter__(self) -> _FakeStreamResponse: + return self._response + + async def __aexit__(self, exc_type, exc, tb) -> bool: + await self._response.aclose() + return False + + +class _FakeAsyncClientFactory: + def __init__(self, responses) -> None: + self._responses = [list(response) for response in responses] + self.stream_calls: list[dict] = [] + self.client_kwargs: list[dict] = [] + + def build(self, *args, **kwargs): + self.client_kwargs.append(dict(kwargs)) + return _FakeAsyncClient(self) + + +class _FakeAsyncClient: + def __init__(self, factory: _FakeAsyncClientFactory) -> None: + self._factory = factory + self.closed = False + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb) -> bool: + await self.aclose() + return False + + def stream(self, method: str, url: str, **kwargs): + if not self._factory._responses: + raise AssertionError("No fake SSE response configured") + + lines = self._factory._responses.pop(0) + self._factory.stream_calls.append( + { + "method": method, + "url": url, + "params": dict(kwargs.get("params") or {}), + "headers": dict(kwargs.get("headers") or {}), + "json": kwargs.get("json"), + } + ) + return _FakeStreamContextManager( + _FakeStreamResponse(lines, method=method, url=url) + ) + + async def aclose(self) -> None: + self.closed = True + + +class _FakeTimeline: + def __init__(self) -> None: + self.events: list[tuple[str, dict]] = [] + + def emit(self, event: str, **fields) -> None: + self.events.append((event, dict(fields))) + + +class RemoteAgentSSEAdapterTests(unittest.IsolatedAsyncioTestCase): + async def test_prepare_captures_session_and_post_reads_own_short_stream(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "ready", + { + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + "session_id": "sessao-123", + "speech_id": "11111111-1111-4111-8111-111111111111", + "is_interruptible": False, + "metadata": {"wait_timeout_seconds": 30}, + }, + ), + *_sse_event("prefetch_done", {"session_id": "sessao-123"}), + ], + [ + *_sse_event( + "ready", + { + "message": "Processando.", + "actions": ["chat", "ping"], + "session_id": "sessao-123", + }, + ), + *_sse_event( + "result", + { + "type": "result", + "action": "chat", + "result": { + "type": "final", + "content": "A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + "speech_id": "7f3a0000-0000-4000-8000-000000000000", + "is_interruptible": True, + "tool_calls": [], + }, + }, + ), + ], + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/sse", + request_context={ + "agent": "conta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0005", + "ANI": "3133338888", + "GSM": "551199993339999", + "ID_FATURA": "fat-555", + "callIdGed": "GED-555", + "protocol_id": "PRT-555", + "sessionId": "call-session-555", + }, + ) + await adapter.prepare(True, "PRT-555") + + ready_reply = await adapter.run("") + ready_message_id = _assert_uuid(self, ready_reply.metadata["message_id"]) + self.assertEqual( + ready_reply, + BackendReply( + stage="PRESENTATION", + text="Posso verificar sua fatura.", + done=False, + export_payload=None, + metadata={ + "message_id": ready_message_id, + "wait_timeout_seconds": 30, + "speech_id": "11111111-1111-4111-8111-111111111111", + "is_interruptible": False, + "agent_message_type": "ready", + "expects_user_response": True, + "drop_user_input_while_speaking": True, + }, + ), + ) + self.assertEqual(len(factory.stream_calls), 1) + + get_call = factory.stream_calls[0] + self.assertEqual(get_call["method"], "GET") + self.assertEqual(get_call["url"], "http://agent.internal/agent/sse") + self.assertEqual( + get_call["params"], + { + "msisdn": "551199993339999", + "invoice_id": "fat-555", + "ani": "3133338888", + "protocol_id": "PRT-555", + "session_id": "call-session-555", + "message_id": ready_message_id, + "channelId": "ura", + "uraCallId": "GED-555", + }, + ) + + await adapter.set_processing_interruption( + listened_text="pode repetir", + skipped=False, + speech_id="11111111-1111-4111-8111-111111111111", + ) + await adapter.inject_idle_nudge("Alo, voce esta ai?") + + reply = await adapter.run("teste") + reply_message_id = _assert_uuid(self, reply.metadata["message_id"]) + self.assertNotEqual(reply_message_id, ready_message_id) + self.assertEqual( + reply, + BackendReply( + stage="PRESENTATION", + text="A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + done=False, + export_payload={ + "type": "final", + "content": "A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + "speech_id": "7f3a0000-0000-4000-8000-000000000000", + "is_interruptible": True, + "tool_calls": [], + }, + metadata={ + "message_id": reply_message_id, + "speech_id": "7f3a0000-0000-4000-8000-000000000000", + "is_interruptible": True, + "agent_message_type": "result", + "agent_result_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + + post_call = factory.stream_calls[1] + self.assertEqual(post_call["method"], "POST") + self.assertEqual(post_call["url"], "http://agent.internal/agent/sse") + self.assertEqual( + post_call["params"], + { + "session_id": "sessao-123", + "ani": "3133338888", + "protocol_id": "PRT-555", + "message_id": reply_message_id, + "channelId": "ura", + "uraCallId": "GED-555", + }, + ) + self.assertNotEqual(post_call["params"]["message_id"], post_call["params"]["session_id"]) + self.assertEqual(adapter._remote_session_id, "sessao-123") + self.assertEqual(post_call["json"]["action"], "chat") + self.assertEqual(post_call["json"]["payload"]["message"], "teste") + self.assertEqual( + post_call["json"]["payload"]["message_id"], + reply_message_id, + ) + self.assertNotIn("msisdn", post_call["json"]["payload"]) + self.assertNotIn("invoice_id", post_call["json"]["payload"]) + self.assertEqual( + post_call["json"]["payload"]["processing_interruption"], + { + "speech_id": "11111111-1111-4111-8111-111111111111", + "heard_text": "pode repetir", + }, + ) + self.assertNotIn("speech_interruption", post_call["json"]["payload"]) + self.assertNotIn("interruption", post_call["json"]["payload"]) + self.assertEqual( + post_call["json"]["payload"]["events"], + [{"type": "idle_nudge", "text": "Alo, voce esta ai?"}], + ) + self.assertFalse(adapter.supports_server_push()) + + async def test_conta_feedback_is_buffered_and_run_waits_for_top_level_final(self) -> None: + feedback_payload = { + "type": "feedback", + "text": "Ainda estou consultando sua fatura.", + } + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "ready", + { + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + "session_id": "sessao-125", + }, + ), + *_sse_event("prefetch_done", {"session_id": "sessao-125"}), + ], + [ + *_sse_event("message", feedback_payload), + *_sse_event( + "message", + { + "type": "final", + "content": "Encontrei os detalhes da sua fatura.", + }, + ), + ], + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/sse", + request_context={ + "agent": "conta", + "GSM": "551199993339999", + "ID_FATURA": "fat-557", + }, + ) + await adapter.prepare(True, "PRT-557") + await adapter.run("") + + reply = await adapter.run("teste") + feedback_reply = await adapter.wait_for_server_push() + reply_message_id = _assert_uuid(self, reply.metadata["message_id"]) + + self.assertEqual( + reply, + BackendReply( + stage="PRESENTATION", + text="Encontrei os detalhes da sua fatura.", + done=False, + export_payload=None, + metadata={ + "message_id": reply_message_id, + "agent_message_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + self.assertEqual( + feedback_reply, + BackendReply( + stage="PRESENTATION", + text="Ainda estou consultando sua fatura.", + done=False, + export_payload=None, + metadata={ + "message_id": reply_message_id, + "agent_message_type": "feedback", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + "event": "feedback", + "payload": feedback_payload, + }, + ), + ) + + async def test_run_marks_expected_agent_result_type_as_done(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "ready", + { + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + "session_id": "sessao-124", + }, + ), + *_sse_event("prefetch_done", {"session_id": "sessao-124"}), + ], + [ + *_sse_event( + "result", + { + "type": "result", + "action": "chat", + "result": { + "type": "erro_no_match", + "content": "Atendimento encerrado por falta de correspondencia.", + "tool_calls": [], + }, + }, + ), + ], + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/sse", + request_context={ + "agent": "conta", + "GSM": "551199993339999", + "ID_FATURA": "fat-556", + }, + ) + await adapter.prepare(True, "PRT-556") + await adapter.run("") + + reply = await adapter.run("teste") + reply_message_id = _assert_uuid(self, reply.metadata["message_id"]) + + self.assertEqual( + reply, + BackendReply( + stage="DONE", + text="Atendimento encerrado por falta de correspondencia.", + done=True, + export_payload={ + "type": "erro_no_match", + "content": "Atendimento encerrado por falta de correspondencia.", + "tool_calls": [], + }, + metadata={ + "message_id": reply_message_id, + "agent_message_type": "result", + "agent_result_type": "erro_no_match", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + }, + ), + ) + + async def test_conta_result_feedback_is_buffered_and_run_waits_for_final(self) -> None: + feedback_result = { + "type": "feedback", + "content": "Perfeito, seguiremos com o cancelamento. Aguarde um instante.", + "speech_id": "332a02d1-031d-4d08-b14a-f896fe630fbc", + "is_interruptible": True, + "tool_calls": [], + "metadata": {"stage": "feedback"}, + } + feedback_payload = { + "type": "result", + "action": "chat", + "result": feedback_result, + } + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "ready", + { + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + "session_id": "sessao-126", + }, + ), + *_sse_event("prefetch_done", {"session_id": "sessao-126"}), + ], + [ + *_sse_event( + "result", + feedback_payload, + ), + *_sse_event( + "result", + { + "type": "result", + "action": "chat", + "result": { + "type": "final", + "content": "Cancelamento concluido.", + "tool_calls": [], + }, + }, + ), + ], + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/sse", + request_context={ + "agent": "conta", + "GSM": "551199993339999", + "ID_FATURA": "fat-558", + }, + ) + await adapter.prepare(True, "PRT-558") + await adapter.run("") + + reply = await adapter.run("correto") + feedback_reply = await adapter.wait_for_server_push() + reply_message_id = _assert_uuid(self, reply.metadata["message_id"]) + + self.assertEqual( + reply, + BackendReply( + stage="PRESENTATION", + text="Cancelamento concluido.", + done=False, + export_payload={ + "type": "final", + "content": "Cancelamento concluido.", + "tool_calls": [], + }, + metadata={ + "message_id": reply_message_id, + "agent_message_type": "result", + "agent_result_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + self.assertEqual( + feedback_reply, + BackendReply( + stage="PRESENTATION", + text="Perfeito, seguiremos com o cancelamento. Aguarde um instante.", + done=False, + export_payload=None, + metadata={ + "message_id": reply_message_id, + "speech_id": "332a02d1-031d-4d08-b14a-f896fe630fbc", + "is_interruptible": False, + "agent_message_type": "result", + "agent_result_type": "feedback", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "event": "feedback", + "payload": feedback_payload, + }, + ), + ) + + async def test_prepare_stops_after_ready_even_if_stream_has_late_error(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "ready", + { + "message": "Posso verificar sua fatura.", + "actions": ["chat"], + "session_id": "sessao-999", + }, + ), + *_sse_event("error", {"error": "late stream error after ready"}), + ] + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/sse", + request_context={ + "agent": "conta", + "GSM": "5511999900000", + "ID_FATURA": "fat-999", + }, + ) + + await adapter.prepare(True, "ignored-protocol") + ready_reply = await adapter.run("") + ready_message_id = _assert_uuid(self, ready_reply.metadata["message_id"]) + + self.assertEqual( + ready_reply, + BackendReply( + stage="PRESENTATION", + text="Posso verificar sua fatura.", + done=False, + export_payload=None, + metadata={ + "message_id": ready_message_id, + "agent_message_type": "ready", + "expects_user_response": True, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + }, + ), + ) + self.assertEqual(adapter._remote_session_id, "sessao-999") + self.assertEqual(len(factory.stream_calls), 1) + self.assertEqual(factory.stream_calls[0]["method"], "GET") + self.assertEqual( + factory.stream_calls[0]["params"], + { + "msisdn": "5511999900000", + "invoice_id": "fat-999", + "message_id": ready_message_id, + "channelId": "ura", + }, + ) + + async def test_run_accumulates_progress_chunks_until_done(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "ready", + { + "message": "Posso ajudar.", + "actions": ["chat"], + "session_id": "PRT-777", + }, + ) + ], + [ + *_sse_event("progress", {"delta": "Oferta "}), + *_sse_event("progress", {"delta": "liberada"}), + *_sse_event("result", {"stage": "DONE"}), + ], + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/sse", + request_context={ + "agent": "conta", + "GSM": "5511999900000", + }, + ) + await adapter.prepare(True, "PRT-777") + await adapter.run("") + + reply = await adapter.run("teste") + reply_message_id = _assert_uuid(self, reply.metadata["message_id"]) + self.assertEqual( + reply, + BackendReply( + stage="DONE", + text="Oferta liberada", + done=True, + export_payload={"stage": "DONE", "type": "result"}, + metadata={ + "message_id": reply_message_id, + "agent_message_type": "result", + "agent_result_type": "result", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + }, + ), + ) + + async def test_end_service_once_skips_end_action_when_backend_does_not_advertise_it(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "ready", + { + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + "session_id": "PRT-559", + }, + ) + ] + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/sse", + request_context={ + "agent": "conta", + "GSM": "5511999900003", + }, + ) + await adapter.prepare(True, "PRT-559") + ready_reply = await adapter.run("") + self.assertEqual(ready_reply.text, "Posso verificar sua fatura.") + + end_reply = await adapter.end_service_once() + self.assertEqual( + end_reply, + BackendReply( + stage="DONE", + text="", + done=True, + export_payload=[], + ), + ) + self.assertEqual(len(factory.stream_calls), 1) + + async def test_oferta_initial_turn_posts_execute_and_streams_sse_events(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "schedule_message", + { + "scheduledMessage": { + "timeInSeconds": 5, + "message": "Aguarde, estou consultando sua oferta.", + } + }, + ), + *_sse_event( + "message", + { + "response": "Tenho uma oferta disponivel para voce.", + "sessionId": "remote-session-001", + "additionalInformations": {"foo": "bar"}, + }, + ), + *_sse_event( + "done", + { + "status": "completed", + "additionalInformations": {"service_status": "RESOLVED"}, + }, + ), + ], + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/execute", + request_context={ + "agent": "oferta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0005", + "ANI": "3133338888", + "GSM": "551199993339999", + "callIdGed": "GED-555", + "protocolo": "PRT-555", + "assetId": "asset-555", + "channelId": "ura", + }, + ) + await adapter.prepare(True, "PRT-555") + + first_reply = await adapter.run("") + message_id = _assert_uuid(self, first_reply.metadata["message_id"]) + self.assertEqual( + first_reply, + BackendReply( + stage="PRESENTATION", + text="Aguarde, estou consultando sua oferta.", + done=False, + export_payload=None, + metadata={ + "message_id": message_id, + "event": "schedule_message", + "scheduledMessage": { + "timeInSeconds": 5, + "message": "Aguarde, estou consultando sua oferta.", + }, + "payload": { + "scheduledMessage": { + "timeInSeconds": 5, + "message": "Aguarde, estou consultando sua oferta.", + }, + "type": "schedule_message", + }, + }, + ), + ) + + message_reply = await adapter.wait_for_server_push() + self.assertEqual( + message_reply, + BackendReply( + stage="PRESENTATION", + text="Tenho uma oferta disponivel para voce.", + done=False, + export_payload=None, + metadata={ + "message_id": message_id, + "event": "message", + "payload": { + "response": "Tenho uma oferta disponivel para voce.", + "sessionId": "remote-session-001", + "additionalInformations": {"foo": "bar"}, + "type": "message", + }, + "sessionId": "remote-session-001", + "additionalInformations": {"foo": "bar"}, + }, + ), + ) + + done_reply = await adapter.wait_for_server_push() + self.assertEqual(done_reply.stage, "DONE") + self.assertTrue(done_reply.done) + self.assertEqual(done_reply.export_payload["type"], "resolvido") + self.assertEqual(adapter._remote_session_id, "remote-session-001") + + self.assertEqual(len(factory.stream_calls), 1) + post_call = factory.stream_calls[0] + self.assertEqual(post_call["method"], "POST") + self.assertEqual(post_call["url"], "http://agent.internal/agent/execute") + self.assertEqual(post_call["params"], {}) + self.assertEqual(post_call["headers"]["Channel-id"], "ura") + self.assertEqual(post_call["json"]["message"], "inicio_atendimento") + self.assertEqual(post_call["json"]["messageId"], message_id) + self.assertEqual( + post_call["json"]["context"], + { + "protocolNumber": "PRT-555", + "protocolo": "PRT-555", + "gsm": "551199993339999", + "uraId": "GED-555", + "callIdGed": "GED-555", + "ani": "3133338888", + "routerCallKey": "0005", + "routerCallKeyDay": "20260329", + "agent": "oferta", + "assetId": "asset-555", + }, + ) + self.assertNotIn("sessionId", post_call["json"]["context"]) + + async def test_oferta_forwards_only_explicit_session_id(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "done", + { + "status": "completed", + "additionalInformations": {"service_status": "UNRESOLVED"}, + }, + ) + ], + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/execute", + request_context={ + "agent": "oferta", + "GSM": "551199993339999", + "callIdGed": "GED-556", + "protocolo": "PRT-556", + "sessionId": "session-556", + }, + ) + await adapter.prepare(True, "PRT-556") + reply = await adapter.run("quero oferta") + + self.assertEqual(reply.stage, "PRESENTATION") + self.assertFalse(reply.done) + self.assertEqual(reply.export_payload["type"], "unresolved") + self.assertEqual( + factory.stream_calls[0]["json"]["context"]["sessionId"], + "session-556", + ) + self.assertEqual(factory.stream_calls[0]["json"]["message"], "quero oferta") + + async def test_oferta_ignores_removed_test_session_id_env(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "done", + { + "status": "completed", + "additionalInformations": {"service_status": "UNRESOLVED"}, + }, + ) + ], + ] + ) + + with mock.patch.dict( + os.environ, + {"REMOTE_AGENT_OFERTA_TEST_SESSION_ID": "dev-test-session"}, + ), mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/execute", + request_context={ + "agent": "oferta", + "GSM": "551199993339999", + "callIdGed": "GED-556", + "protocolo": "PRT-556", + }, + ) + await adapter.prepare(True, "PRT-556") + await adapter.run("quero oferta") + + self.assertNotIn("sessionId", factory.stream_calls[0]["json"]["context"]) + + async def test_oferta_can_disable_tls_verification_by_env(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "done", + { + "status": "completed", + "additionalInformations": {"service_status": "UNRESOLVED"}, + }, + ) + ], + ] + ) + + with mock.patch.dict( + os.environ, + {"REMOTE_AGENT_SSE_TLS_VERIFY_OFERTA": "0"}, + ), mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/execute", + request_context={ + "agent": "oferta", + "GSM": "551199993339999", + "callIdGed": "GED-556", + "protocolo": "PRT-556", + }, + ) + await adapter.prepare(True, "PRT-556") + await adapter.run("quero oferta") + + self.assertEqual(factory.client_kwargs[0]["verify"], False) + + async def test_oferta_timeline_records_sse_event_without_emit_collision(self) -> None: + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "schedule_message", + { + "scheduledMessage": { + "timeInSeconds": 5, + "message": "Processando sua solicitação.", + } + }, + ), + *_sse_event( + "done", + { + "status": "completed", + "additionalInformations": {"service_status": "UNRESOLVED"}, + }, + ), + ], + ] + ) + timeline = _FakeTimeline() + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/execute", + request_context={ + "agent": "oferta", + "GSM": "551199993339999", + "callIdGed": "GED-556", + "protocolo": "PRT-556", + }, + timeline=timeline, + ) + await adapter.prepare(True, "PRT-556") + await adapter.run("") + + response_events = [ + fields + for event_name, fields in timeline.events + if event_name == "remote_agent_response" + ] + self.assertEqual(response_events[0]["sse_event"], "schedule_message") + self.assertNotIn("event", response_events[0]) + + async def test_oferta_done_maps_supported_service_statuses_and_transfer(self) -> None: + cases = [ + ("RESOLVED", "completed", "resolvido", "DONE", True), + ("UNRESOLVED", "completed", "unresolved", "PRESENTATION", False), + ("RESOLVED_WITH_NEW_REQUEST", "completed", "resolved_with_new_request", "PRESENTATION", False), + ("", "completed", "completed", "PRESENTATION", False), + ("RESOLVED", "transferred", "transferred", "PRESENTATION", False), + ] + + for service_status, status, expected_result_type, expected_stage, expected_done in cases: + with self.subTest(service_status=service_status, status=status): + factory = _FakeAsyncClientFactory( + responses=[ + [ + *_sse_event( + "done", + { + "status": status, + "additionalInformations": { + "service_status": service_status, + "changeAgentId": "ATH-001", + "context": "handover", + }, + }, + ) + ], + ] + ) + + with mock.patch( + "app.livekit.adapters.remote_agent_sse_adapter.httpx.AsyncClient", + side_effect=factory.build, + ): + adapter = RemoteAgentSSEAdapter( + intro="oi", + url="http://agent.internal/agent/execute", + request_context={ + "agent": "oferta", + "GSM": "551199993339999", + "callIdGed": "GED-557", + "protocolo": "PRT-557", + }, + ) + await adapter.prepare(True, "PRT-557") + reply = await adapter.run("teste") + + self.assertEqual(reply.stage, expected_stage) + self.assertEqual(reply.done, expected_done) + self.assertEqual(reply.export_payload["type"], expected_result_type) + self.assertEqual(reply.export_payload["status"], status) + self.assertEqual( + reply.export_payload["additionalInformations"]["changeAgentId"], + "ATH-001", + ) + + async def test_idle_nudge_keeps_only_the_last_phrase_before_the_turn(self) -> None: + adapter = RemoteAgentSSEAdapter( + intro="", + url="http://agent.internal/agent/sse", + request_context={"agent": "conta"}, + ) + + # metadata.wait_retry_messages manda tres frases distintas por chamada; + # sem a colapsagem elas viram tres falas do agente no historico remoto. + for retry_text in ( + "Voce ainda esta na linha?", + "Sigo por aqui.", + "Vou confirmar mais uma vez, senao vou precisar encerrar o atendimento.", + ): + await adapter.inject_idle_nudge(retry_text) + + self.assertEqual( + adapter._pending_events, + [ + { + "type": "idle_nudge", + "text": "Vou confirmar mais uma vez, senao vou precisar encerrar o atendimento.", + } + ], + ) + + payload = adapter._build_turn_payload("estou na linha", message_id="MSG-1") + self.assertEqual(len(payload["payload"]["events"]), 1) + + +class BackendFactoryTests(unittest.TestCase): + def test_selects_remote_sse_backend(self) -> None: + with mock.patch.dict( + os.environ, + { + "AGENT_BACKEND": "remote_sse", + "REMOTE_AGENT_SSE_URL_CONTA": "http://agent.internal/conta/sse", + }, + clear=False, + ): + backend = build_agent_backend( + intro="oi", + remote_agent_context={"agent": "conta"}, + streaming=False, + ) + + self.assertIsInstance(backend, RemoteAgentSSEAdapter) + self.assertEqual(backend._url, "http://agent.internal/conta/sse") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/livekit/test_remote_agent_ws_adapter.py b/tests/livekit/test_remote_agent_ws_adapter.py new file mode 100644 index 0000000..981fc93 --- /dev/null +++ b/tests/livekit/test_remote_agent_ws_adapter.py @@ -0,0 +1,1078 @@ +from __future__ import annotations + +import asyncio +import json +import os +import sys +import types +import unittest +import uuid +from unittest import mock +from urllib.parse import parse_qs, urlsplit + +from app.livekit.adapters.agent_backend import BackendReply +from app.livekit.adapters.backend_factory import build_agent_backend +from app.livekit.adapters.fake_remote_ws_adapter import FakeRemoteWSAdapter +from app.livekit.adapters.remote_agent_ws_adapter import RemoteAgentWSAdapter + + +def _assert_uuid(testcase: unittest.TestCase, value: object) -> str: + text = str(value or "") + testcase.assertEqual(str(uuid.UUID(text)), text) + return text + + +class _FakeWebSocket: + def __init__(self, responses, send_responses=None) -> None: + self._responses = list(responses) + self._send_responses = [list(batch) for batch in (send_responses or [])] + self._ready = asyncio.Event() + if self._responses: + self._ready.set() + self.sent_messages = [] + + async def send(self, message: str) -> None: + self.sent_messages.append(message) + if self._send_responses: + self._responses.extend(self._send_responses.pop(0)) + self._ready.set() + + async def recv(self): + while not self._responses: + self._ready.clear() + await self._ready.wait() + return self._responses.pop(0) + + +class _FakeConnect: + def __init__(self, websocket: _FakeWebSocket) -> None: + self._websocket = websocket + + async def __aenter__(self) -> _FakeWebSocket: + return self._websocket + + async def __aexit__(self, exc_type, exc, tb) -> bool: + return False + + +class _FakeWebsocketsModule(types.SimpleNamespace): + def __init__(self, sessions) -> None: + super().__init__() + self._sessions = [] + for spec in sessions: + if isinstance(spec, dict): + initial = spec.get("initial") or [] + send_responses = spec.get("after_sends") or [] + self._sessions.append(_FakeWebSocket(initial, send_responses)) + else: + self._sessions.append(_FakeWebSocket(spec)) + self.calls = [] + + def connect(self, url: str, **kwargs): + if not self._sessions: + raise AssertionError("No fake websocket session configured") + websocket = self._sessions.pop(0) + self.calls.append( + { + "url": url, + "kwargs": kwargs, + "websocket": websocket, + } + ) + return _FakeConnect(websocket) + + +class RemoteAgentWSAdapterTests(unittest.IsolatedAsyncioTestCase): + async def test_run_sends_turn_payload_and_parses_response_for_non_conta_agents(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + [ + json.dumps( + { + "type": "final", + "stage": "ARGUMENTATION", + "text": "Tenho uma oferta para voce.", + } + ) + ], + [ + json.dumps( + { + "type": "final", + "stage": "FORMALIZATION", + "text": "Continuando sem interrupcao.", + } + ) + ], + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/turn", + url_by_agent={ + "oferta": "ws://agent.internal/oferta", + "cobranca": "ws://agent.internal/cobranca", + }, + request_context={ + "agent": "oferta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0001", + "ANI": "3133334444", + "GSM": "31999999999", + "callIdGed": "GED-123", + }, + ) + await adapter.prepare(True, "PRT-123") + await adapter.set_interruption( + True, + listened_text="pode repetir", + skipped=False, + speech_id="7f3a0000-0000-4000-8000-000000000000", + ) + await adapter.inject_idle_nudge("Alo, voce esta ai?") + + reply = await adapter.run({"text": "quero detalhes"}) + self.assertEqual( + reply, + BackendReply( + stage="ARGUMENTATION", + text="Tenho uma oferta para voce.", + done=False, + export_payload=None, + ), + ) + + first_payload = json.loads(fake_ws_module.calls[0]["websocket"].sent_messages[0]) + self.assertEqual(fake_ws_module.calls[0]["url"], "ws://agent.internal/oferta") + self.assertIn("timestamp", first_payload) + self.assertEqual(first_payload["agent"], "oferta") + self.assertEqual(first_payload["RouterCallKeyDay"], "20260329") + self.assertEqual(first_payload["RouterCallKey"], "0001") + self.assertEqual(first_payload["ANI"], "3133334444") + self.assertEqual(first_payload["GSM"], "31999999999") + self.assertEqual(first_payload["callIdGed"], "GED-123") + self.assertEqual(first_payload["protocol"], "PRT-123") + self.assertEqual(first_payload["text"], "quero detalhes") + self.assertEqual( + first_payload["speech_interruption"], + { + "speech_id": "7f3a0000-0000-4000-8000-000000000000", + "heard_text": "pode repetir", + }, + ) + self.assertNotIn("processing_interruption", first_payload) + self.assertEqual( + first_payload["events"], + [{"type": "idle_nudge", "text": "Alo, voce esta ai?"}], + ) + + reply_2 = await adapter.run({"text": "segue"}) + self.assertEqual( + reply_2, + BackendReply( + stage="FORMALIZATION", + text="Continuando sem interrupcao.", + done=False, + export_payload=None, + ), + ) + + second_payload = json.loads(fake_ws_module.calls[1]["websocket"].sent_messages[0]) + self.assertNotIn("interruption", second_payload) + self.assertNotIn("speech_interruption", second_payload) + self.assertNotIn("events", second_payload) + + async def test_conta_prepare_buffers_ready_and_reuses_socket_for_proactive_results(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + { + "initial": [ + json.dumps( + { + "type": "ready", + "message": "Entendi que voce precisa de mais detalhes sobre os valores da sua fatura.", + "actions": ["chat", "ping"], + "speech_id": "11111111-1111-4111-8111-111111111111", + "is_interruptible": False, + "metadata": {"wait_timeout_seconds": 30}, + } + ) + ], + "after_sends": [ + [ + json.dumps( + { + "type": "proactive_result", + "action": "explicar_fatura", + "result": { + "type": "final", + "content": "A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + "speech_id": "7f3a0000-0000-4000-8000-000000000000", + "is_interruptible": True, + "tool_calls": [], + }, + } + ), + json.dumps( + { + "type": "proactive_result", + "action": "explicar_fatura", + "result": { + "type": "final", + "content": "Esses sao os principais fatores que podem gerar a diferenca na fatura.", + "tool_calls": [], + }, + } + ), + ] + ], + } + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/conta", + request_context={ + "agent": "conta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0005", + "ANI": "3133338888", + "GSM": "551199993339999", + "ID_FATURA": "fat-555", + "callIdGed": "GED-555", + "protocol_id": "PRT-555", + "sessionId": "call-session-555", + }, + ) + await adapter.prepare(True, "PRT-555") + + ready_reply = await adapter.run("") + connection_message_id = _assert_uuid(self, ready_reply.metadata["message_id"]) + + self.assertEqual( + ready_reply, + BackendReply( + stage="PRESENTATION", + text="Entendi que voce precisa de mais detalhes sobre os valores da sua fatura.", + done=False, + export_payload=None, + metadata={ + "message_id": connection_message_id, + "wait_timeout_seconds": 30, + "speech_id": "11111111-1111-4111-8111-111111111111", + "is_interruptible": False, + "agent_message_type": "ready", + "expects_user_response": True, + "drop_user_input_while_speaking": True, + }, + ), + ) + + ws = fake_ws_module.calls[0]["websocket"] + self.assertEqual(ws.sent_messages, []) + + await adapter.set_interruption( + True, + listened_text="pode repetir", + skipped=False, + speech_id="11111111-1111-4111-8111-111111111111", + ) + await adapter.inject_idle_nudge("Alo, voce esta ai?") + + reply = await adapter.run("teste") + self.assertEqual(reply.metadata["message_id"], connection_message_id) + + self.assertEqual( + reply, + BackendReply( + stage="PRESENTATION", + text="A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + done=False, + export_payload={ + "type": "final", + "content": "A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + "speech_id": "7f3a0000-0000-4000-8000-000000000000", + "is_interruptible": True, + "tool_calls": [], + }, + metadata={ + "message_id": connection_message_id, + "speech_id": "7f3a0000-0000-4000-8000-000000000000", + "is_interruptible": True, + "agent_message_type": "proactive_result", + "agent_result_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + + first_payload = json.loads(ws.sent_messages[0]) + parsed_url = urlsplit(fake_ws_module.calls[0]["url"]) + self.assertEqual( + f"{parsed_url.scheme}://{parsed_url.netloc}{parsed_url.path}", + "ws://agent.internal/conta", + ) + self.assertEqual( + parse_qs(parsed_url.query), + { + "msisdn": ["551199993339999"], + "current_invoice_number": ["fat-555"], + "ani": ["3133338888"], + "protocol_id": ["PRT-555"], + "session_id": ["call-session-555"], + "message_id": [connection_message_id], + "channelId": ["ura"], + "uraCallId": ["GED-555"], + }, + ) + self.assertEqual(first_payload["action"], "chat") + self.assertEqual(first_payload["payload"]["message"], "teste") + self.assertEqual(first_payload["payload"]["channel"], "SUPERVISOR") + self.assertEqual(first_payload["payload"]["msisdn"], "551199993339999") + self.assertEqual(first_payload["payload"]["current_invoice_number"], "fat-555") + self.assertEqual( + first_payload["payload"]["speech_interruption"], + { + "speech_id": "11111111-1111-4111-8111-111111111111", + "heard_text": "pode repetir", + }, + ) + self.assertNotIn("processing_interruption", first_payload["payload"]) + self.assertNotIn("interruption", first_payload["payload"]) + self.assertEqual( + first_payload["payload"]["events"], + [{"type": "idle_nudge", "text": "Alo, voce esta ai?"}], + ) + + reply_2 = await adapter.wait_for_server_push() + self.assertEqual( + reply_2, + BackendReply( + stage="PRESENTATION", + text="Esses sao os principais fatores que podem gerar a diferenca na fatura.", + done=False, + export_payload={ + "type": "final", + "content": "Esses sao os principais fatores que podem gerar a diferenca na fatura.", + "tool_calls": [], + }, + metadata={ + "message_id": connection_message_id, + "agent_message_type": "proactive_result", + "agent_result_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + self.assertEqual(len(ws.sent_messages), 1) + + async def test_conta_run_parses_top_level_final_content(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + { + "initial": [ + json.dumps( + { + "type": "ready", + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + } + ) + ], + "after_sends": [ + [ + json.dumps( + { + "type": "final", + "content": "Infelizmente **não encontrei** ", + } + ) + ] + ], + } + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/conta", + request_context={ + "agent": "conta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0006", + "ANI": "3133330000", + "GSM": "5511999900000", + "callIdGed": "GED-556", + }, + ) + await adapter.prepare(True, "PRT-556") + + ready_reply = await adapter.run("") + self.assertEqual(ready_reply.text, "Posso verificar sua fatura.") + + reply = await adapter.run("teste") + connection_message_id = _assert_uuid(self, reply.metadata["message_id"]) + + self.assertEqual( + reply, + BackendReply( + stage="PRESENTATION", + text="Infelizmente **não encontrei**", + done=False, + export_payload=None, + metadata={ + "message_id": connection_message_id, + "agent_message_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + + async def test_conta_feedback_is_buffered_and_run_waits_for_top_level_final(self) -> None: + feedback_payload = { + "type": "feedback", + "text": "Ainda estou consultando sua fatura.", + } + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + { + "initial": [ + json.dumps( + { + "type": "ready", + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + } + ) + ], + "after_sends": [ + [ + json.dumps(feedback_payload), + json.dumps( + { + "type": "final", + "content": "Encontrei os detalhes da sua fatura.", + } + ), + ] + ], + } + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/conta", + request_context={ + "agent": "conta", + "GSM": "5511999900000", + "callIdGed": "GED-556", + }, + ) + await adapter.prepare(True, "PRT-556") + + ready_reply = await adapter.run("") + self.assertEqual(ready_reply.text, "Posso verificar sua fatura.") + + reply = await adapter.run("teste") + feedback_reply = await adapter.wait_for_server_push() + connection_message_id = _assert_uuid(self, reply.metadata["message_id"]) + + self.assertEqual( + reply, + BackendReply( + stage="PRESENTATION", + text="Encontrei os detalhes da sua fatura.", + done=False, + export_payload=None, + metadata={ + "message_id": connection_message_id, + "agent_message_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + self.assertEqual( + feedback_reply, + BackendReply( + stage="PRESENTATION", + text="Ainda estou consultando sua fatura.", + done=False, + export_payload=None, + metadata={ + "message_id": connection_message_id, + "agent_message_type": "feedback", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + "event": "feedback", + "payload": feedback_payload, + }, + ), + ) + + async def test_conta_run_marks_expected_agent_result_type_as_done(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + { + "initial": [ + json.dumps( + { + "type": "ready", + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + } + ) + ], + "after_sends": [ + [ + json.dumps( + { + "type": "result", + "action": "chat", + "result": { + "type": "resolvido", + "content": "Atendimento encerrado com sucesso.", + "tool_calls": [ + { + "id": "call-final", + "name": "finalizar_atendimento", + "args": { + "status": "resolvido", + "summary": "Resumo da conversa.", + }, + "result": { + "success": True, + "status": "resolvido", + "summary": "Encerramento realizado pelo agente.", + }, + } + ], + }, + } + ) + ] + ], + } + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/conta", + request_context={ + "agent": "conta", + "GSM": "5511999900000", + "callIdGed": "GED-560", + }, + ) + await adapter.prepare(True, "PRT-560") + + ready_reply = await adapter.run("") + self.assertEqual(ready_reply.text, "Posso verificar sua fatura.") + + reply = await adapter.run("pode encerrar") + connection_message_id = _assert_uuid(self, reply.metadata["message_id"]) + + self.assertEqual( + reply, + BackendReply( + stage="DONE", + text="Atendimento encerrado com sucesso.", + done=True, + export_payload={ + "type": "resolvido", + "content": "Atendimento encerrado com sucesso.", + "tool_calls": [ + { + "id": "call-final", + "name": "finalizar_atendimento", + "args": { + "status": "resolvido", + "summary": "Resumo da conversa.", + }, + "result": { + "success": True, + "status": "resolvido", + "summary": "Encerramento realizado pelo agente.", + }, + } + ], + }, + metadata={ + "message_id": connection_message_id, + "agent_message_type": "result", + "agent_result_type": "resolvido", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + }, + ), + ) + + async def test_conta_result_feedback_is_buffered_and_run_waits_for_final(self) -> None: + feedback_result = { + "type": "feedback", + "content": "Perfeito, seguiremos com o cancelamento. Aguarde um instante.", + "speech_id": "332a02d1-031d-4d08-b14a-f896fe630fbc", + "is_interruptible": True, + "tool_calls": [], + "metadata": {"stage": "feedback"}, + } + feedback_payload = { + "type": "result", + "action": "chat", + "result": feedback_result, + } + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + { + "initial": [ + json.dumps( + { + "type": "ready", + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + } + ) + ], + "after_sends": [ + [ + json.dumps(feedback_payload), + json.dumps( + { + "type": "result", + "action": "chat", + "result": { + "type": "final", + "content": "Cancelamento concluido.", + "tool_calls": [], + }, + } + ), + ] + ], + } + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/conta", + request_context={ + "agent": "conta", + "GSM": "5511999900000", + "callIdGed": "GED-561", + }, + ) + await adapter.prepare(True, "PRT-561") + + ready_reply = await adapter.run("") + self.assertEqual(ready_reply.text, "Posso verificar sua fatura.") + + reply = await adapter.run("correto") + feedback_reply = await adapter.wait_for_server_push() + connection_message_id = _assert_uuid(self, reply.metadata["message_id"]) + + self.assertEqual( + reply, + BackendReply( + stage="PRESENTATION", + text="Cancelamento concluido.", + done=False, + export_payload={ + "type": "final", + "content": "Cancelamento concluido.", + "tool_calls": [], + }, + metadata={ + "message_id": connection_message_id, + "agent_message_type": "result", + "agent_result_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + self.assertEqual( + feedback_reply, + BackendReply( + stage="PRESENTATION", + text="Perfeito, seguiremos com o cancelamento. Aguarde um instante.", + done=False, + export_payload=None, + metadata={ + "message_id": connection_message_id, + "speech_id": "332a02d1-031d-4d08-b14a-f896fe630fbc", + "is_interruptible": False, + "agent_message_type": "result", + "agent_result_type": "feedback", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "event": "feedback", + "payload": feedback_payload, + }, + ), + ) + + async def test_conta_wait_for_server_push_receives_proactive_result_without_new_run(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + [ + json.dumps( + { + "type": "ready", + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + } + ), + json.dumps( + { + "type": "proactive_result", + "action": "explicar_fatura", + "result": { + "type": "final", + "content": "A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + "tool_calls": [], + }, + } + ), + ] + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/conta", + request_context={ + "agent": "conta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0007", + "ANI": "3133330001", + "GSM": "5511999900001", + "callIdGed": "GED-557", + }, + ) + await adapter.prepare(True, "PRT-557") + + ready_reply = await adapter.run("") + self.assertEqual(ready_reply.text, "Posso verificar sua fatura.") + + push_reply = await adapter.wait_for_server_push() + connection_message_id = _assert_uuid(self, push_reply.metadata["message_id"]) + self.assertEqual( + push_reply, + BackendReply( + stage="PRESENTATION", + text="A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + done=False, + export_payload={ + "type": "final", + "content": "A variacao entre a fatura atual e a anterior pode ocorrer por alguns motivos comuns.", + "tool_calls": [], + }, + metadata={ + "message_id": connection_message_id, + "agent_message_type": "proactive_result", + "agent_result_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + ) + + async def test_conta_can_send_new_chat_after_idle_period_without_killing_reader(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + { + "initial": [ + json.dumps( + { + "type": "ready", + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + } + ) + ], + "after_sends": [ + [ + json.dumps( + { + "type": "proactive_result", + "action": "explicar_fatura", + "result": { + "type": "final", + "content": "Resumo inicial da fatura.", + "tool_calls": [], + }, + } + ) + ], + [ + json.dumps( + { + "type": "final", + "content": "Os servicos ativos sao Netflix e Deezer.", + } + ) + ], + ], + } + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/conta", + request_context={ + "agent": "conta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0008", + "ANI": "3133330002", + "GSM": "5511999900002", + "callIdGed": "GED-558", + }, + read_timeout_s=0.1, + ) + await adapter.prepare(True, "PRT-558") + + ready_reply = await adapter.run("") + self.assertEqual(ready_reply.text, "Posso verificar sua fatura.") + + first_reply = await adapter.run("quero detalhes") + self.assertEqual(first_reply.text, "Resumo inicial da fatura.") + + await asyncio.sleep(0.12) + + second_reply = await adapter.run("Quais servicos eu tenho?") + self.assertEqual(second_reply.text, "Os servicos ativos sao Netflix e Deezer.") + + async def test_conta_end_service_once_skips_end_action_when_backend_does_not_advertise_it(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + [ + json.dumps( + { + "type": "ready", + "message": "Posso verificar sua fatura.", + "actions": ["chat", "ping"], + } + ) + ] + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/conta", + request_context={ + "agent": "conta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0009", + "ANI": "3133330003", + "GSM": "5511999900003", + "callIdGed": "GED-559", + }, + ) + await adapter.prepare(True, "PRT-559") + + ready_reply = await adapter.run("") + self.assertEqual(ready_reply.text, "Posso verificar sua fatura.") + + end_reply = await adapter.end_service_once() + self.assertEqual( + end_reply, + BackendReply( + stage="DONE", + text="", + done=True, + export_payload=[], + ), + ) + + ws = fake_ws_module.calls[0]["websocket"] + self.assertEqual(ws.sent_messages, []) + + async def test_run_accumulates_streamed_chunks_until_done(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + [ + json.dumps({"type": "delta", "delta": "Oferta "}), + json.dumps({"type": "delta", "delta": "liberada"}), + json.dumps({"type": "done", "stage": "DONE"}), + ] + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/turn", + url_by_agent={"cobranca": "ws://agent.internal/cobranca"}, + request_context={ + "agent": "cobra", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0002", + "ANI": "3133335555", + "GSM": "31988887777", + "callIdGed": "GED-456", + }, + ) + await adapter.prepare(True, "PRT-456") + + reply = await adapter.run("teste") + + self.assertEqual( + reply, + BackendReply( + stage="DONE", + text="Oferta liberada", + done=True, + export_payload=None, + ), + ) + payload = json.loads(fake_ws_module.calls[0]["websocket"].sent_messages[0]) + self.assertEqual(fake_ws_module.calls[0]["url"], "ws://agent.internal/cobranca") + self.assertEqual(payload["agent"], "cobranca") + self.assertEqual(payload["callIdGed"], "GED-456") + self.assertNotIn("ID_FATURA", payload) + + async def test_run_uses_default_stage_when_remote_does_not_return_stage(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + [ + json.dumps({"type": "final", "text": "Resposta sem stage."}), + ] + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/turn", + default_stage="presentation", + request_context={ + "agent": "ofert", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0003", + "ANI": "3133336666", + "GSM": "31977776666", + "callIdGed": "GED-777", + }, + ) + await adapter.prepare(True, "PRT-777") + + reply = await adapter.run("teste") + + self.assertEqual( + reply, + BackendReply( + stage="PRESENTATION", + text="Resposta sem stage.", + done=False, + export_payload=None, + ), + ) + + async def test_end_service_once_returns_remote_result_and_is_idempotent(self) -> None: + fake_ws_module = _FakeWebsocketsModule( + sessions=[ + [ + json.dumps( + { + "type": "done", + "result": [{"success_purchase": 1, "protocol": "PRT-999"}], + } + ) + ] + ] + ) + + with mock.patch.dict(sys.modules, {"websockets": fake_ws_module}): + adapter = RemoteAgentWSAdapter( + intro="oi", + url="ws://agent.internal/end", + request_context={ + "agent": "oferta", + "RouterCallKeyDay": "20260329", + "RouterCallKey": "0004", + "ANI": "3133337777", + "GSM": "31966665555", + "callIdGed": "GED-999", + }, + ) + await adapter.prepare(True, "PRT-999") + + output_1 = await adapter.end_service_once() + output_2 = await adapter.end_service_once() + + self.assertEqual( + output_1, + BackendReply( + stage="DONE", + text="", + done=True, + export_payload=[{"success_purchase": 1, "protocol": "PRT-999"}], + ), + ) + self.assertEqual(output_2, output_1) + self.assertEqual(len(fake_ws_module.calls), 1) + + +class BackendFactoryTests(unittest.TestCase): + def test_selects_remote_ws_backend(self) -> None: + with mock.patch.dict( + os.environ, + { + "AGENT_BACKEND": "remote_ws", + "REMOTE_AGENT_WS_URL": "ws://agent.internal/turn", + "REMOTE_AGENT_WS_URL_CONTA": "ws://agent.internal/conta", + }, + clear=False, + ): + backend = build_agent_backend( + intro="oi", + remote_agent_context={"agent": "conta"}, + streaming=False, + ) + + self.assertIsInstance(backend, RemoteAgentWSAdapter) + self.assertEqual(backend._url, "ws://agent.internal/turn") + + def test_selects_fake_remote_ws_backend(self) -> None: + with mock.patch.dict( + os.environ, + { + "REMOTE_AGENT_WS_URL": "ws://agent.internal/real", + "REMOTE_AGENT_WS_FAKE_URL": "ws://127.0.0.1:8000/fake-agent/ws", + }, + clear=False, + ): + backend = build_agent_backend( + intro="oi", + backend_name="remote_ws_fake", + remote_agent_context={"agent": "conta"}, + streaming=False, + ) + + self.assertIsInstance(backend, FakeRemoteWSAdapter) + self.assertEqual(backend._backend_label, "remote_ws_fake") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/livekit/test_runtime.py b/tests/livekit/test_runtime.py new file mode 100644 index 0000000..5cd6dd6 --- /dev/null +++ b/tests/livekit/test_runtime.py @@ -0,0 +1,4802 @@ +from __future__ import annotations + +import asyncio +import importlib +import json +import os +import sys +import time +import types +import unittest +import uuid +from pathlib import Path +from tempfile import TemporaryDirectory +from typing import Any +from unittest import mock +from types import SimpleNamespace + + +def _install_fake_livekit() -> None: + try: + importlib.import_module("livekit.agents") + return + except ImportError: + pass + + livekit_module = types.ModuleType("livekit") + agents_module = types.ModuleType("livekit.agents") + + class UserInputTranscribedEvent: + def __init__(self, transcript: str = "", is_final: bool = True) -> None: + self.transcript = transcript + self.is_final = is_final + + class _AudioInputOptions: + def __init__(self, **kwargs) -> None: + self.kwargs = kwargs + + class _AudioOutputOptions: + def __init__(self, **kwargs) -> None: + self.kwargs = kwargs + + class _RoomOptions: + def __init__(self, **kwargs) -> None: + self.kwargs = kwargs + + agents_module.UserInputTranscribedEvent = UserInputTranscribedEvent + agents_module.room_io = SimpleNamespace( + AudioInputOptions=_AudioInputOptions, + AudioOutputOptions=_AudioOutputOptions, + RoomOptions=_RoomOptions, + ) + livekit_module.agents = agents_module + + sys.modules["livekit"] = livekit_module + sys.modules["livekit.agents"] = agents_module + + +_install_fake_livekit() + +from app.livekit.policies.finalization_policy import FinalizationPolicy +from app.livekit.policies.idle_policy import IdlePolicy, IdlePolicyConfig +from app.livekit.policies.interrupt_policy import InterruptPolicy +from app.livekit.adapters.speech_service import SpeechService +from app.livekit.adapters.agent_backend import BackendReply +from app.livekit.runtime.call_runtime import ( + CallRuntime, + InflightBackendWaitNoticeResult, + InflightBackendWaitTimedOut, +) +from app.livekit.runtime.command_executor import RuntimeCommandExecutor +from app.livekit.runtime.commands import ( + EndServiceOnce, + StartSpeech, + ExportSession, + ExtractSpokenText, + InjectIdleNudge, + InterruptSpeech, + NotifyBridgeDone, + NotifyBridgeStop, + RunPipelineInput, + SetPendingInterrupt, + SetPipelineInterruption, + SetPipelineProcessingInterruption, + StartSession, +) +from app.livekit.runtime.scheduler import TimerScheduler +from app.livekit.runtime.state import RuntimeConfig +from app.livekit.vad_dynamic_threshold import ( + DynamicVADThresholdConfig, + DynamicVADThresholdController, +) +from app.providers.stt_internal_livekit import stt_text_with_single_word_threshold +from app.utils.logging import StructuredLogContext +from app.utils.turn_ids import ( + peek_started_turn_message_id, + register_started_turn_message_id, + register_transcribed_turn, + reset_turn_message_sequence, +) + +WAIT_LONG_AUDIO_PATH = str( + Path(__file__).resolve().parents[2] + / "src" + / "app" + / "livekit" + / "assets" + / "comfort" + / "long" + / "01.wav" +) +DEFAULT_WAIT_TEXT = "Um momento, ainda estou consultando para te ajudar." + + +def _expected_wait_text_for_audio(path: str | Path, fallback: str = DEFAULT_WAIT_TEXT) -> str: + text_path = Path(path).with_suffix(".txt") + if not text_path.is_file(): + return fallback + + for encoding in ("utf-8-sig", "utf-8", "latin-1"): + try: + text = " ".join(text_path.read_text(encoding=encoding).split()) + except UnicodeDecodeError: + continue + except OSError: + return fallback + return text or fallback + + return fallback + + +class _NullLogger: + def info(self, *args, **kwargs) -> None: + return None + + def debug(self, *args, **kwargs) -> None: + return None + + def warning(self, *args, **kwargs) -> None: + return None + + def exception(self, *args, **kwargs) -> None: + return None + + +class _AsyncTestCase(unittest.IsolatedAsyncioTestCase): + async def asyncSetUp(self) -> None: + asyncio.get_running_loop().set_debug(False) + + +class _FakeAgent: + def __init__(self) -> None: + self._ready = asyncio.Event() + self._ready.set() + self._run_lock = asyncio.Lock() + self.pipeline = object() + self.pending_interrupts: list[tuple[str, bool, str]] = [] + self._consumed_interrupt = (None, "", False) + self.end_calls = 0 + + async def set_pending_interrupt( + self, + *, + listened_text: str = "", + skipped: bool = False, + speech_id: str = "", + ) -> None: + self.pending_interrupts.append((listened_text, skipped, speech_id)) + + async def consume_pending_interrupt(self): + return self._consumed_interrupt + + async def end_service_once(self): + self.end_calls += 1 + return BackendReply(stage="DONE", done=True, export_payload=[{"status": "ok"}]) + + +class _FakePushPipeline: + def __init__(self, replies) -> None: + self._replies = asyncio.Queue() + for reply in replies: + self._replies.put_nowait(reply) + + def supports_server_push(self) -> bool: + return True + + async def wait_for_server_push(self): + if self._replies.empty(): + return None + return await self._replies.get() + + +class _FakeInflightPushPipeline: + def __init__(self) -> None: + self._replies = asyncio.Queue() + self._closed = asyncio.Event() + + def supports_inflight_backend_push(self) -> bool: + return True + + def supports_server_push(self) -> bool: + return False + + async def push(self, reply) -> None: + await self._replies.put(reply) + + def close(self) -> None: + self._closed.set() + + async def wait_for_server_push(self): + if not self._replies.empty(): + return await self._replies.get() + if self._closed.is_set(): + return None + + get_task = asyncio.create_task(self._replies.get()) + close_task = asyncio.create_task(self._closed.wait()) + done, _pending = await asyncio.wait( + {get_task, close_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + if get_task in done: + close_task.cancel() + await asyncio.gather(close_task, return_exceptions=True) + return get_task.result() + + get_task.cancel() + await asyncio.gather(get_task, return_exceptions=True) + return None + + +class _FakeAgentInput: + def __init__(self) -> None: + self.audio_enabled = True + self.calls = [] + + def set_audio_enabled(self, enabled: bool) -> None: + self.audio_enabled = bool(enabled) + self.calls.append(bool(enabled)) + + +class _FakeSession: + def __init__(self) -> None: + self.handlers = {} + self.start_calls = [] + self.input = _FakeAgentInput() + + def on(self, event_name: str): + def _decorator(fn): + self.handlers[event_name] = fn + return fn + + return _decorator + + async def start(self, **kwargs) -> None: + self.start_calls.append(kwargs) + + +class _FakeRoom: + def __init__(self) -> None: + self.name = "room-test" + self.remote_participants = {} + self.handlers = {} + + def on(self, event_name: str): + def _decorator(fn): + self.handlers[event_name] = fn + return fn + + return _decorator + + +class _FakeTimeline: + def __init__(self) -> None: + self.events = [] + + def emit(self, event: str, **fields) -> None: + self.events.append((event, fields)) + + +class _FakeContext: + def __init__(self) -> None: + self.room = _FakeRoom() + self.shutdown_callbacks = [] + + def add_shutdown_callback(self, callback) -> None: + self.shutdown_callbacks.append(callback) + + +class _FakeCommandExecutor: + def __init__(self) -> None: + self.commands = [] + self.pipeline_result = BackendReply(stage="PRESENTATION", text="resposta") + self.end_reply = BackendReply(stage="DONE", done=True, export_payload=[{"status": "ok"}]) + self.spoken_text = "trecho falado" + self.speech_handle = SimpleNamespace( + id="speech-test", + interrupted=False, + ) + self.speech_handles = None + self.run_pipeline_side_effect = None + self.on_start_session: Any = None + self.wait_for_playout_delay_s = 0.0 + self.wait_for_playout_delays_s = [] + self.wait_for_playout_exceptions = [] + self.wait_for_playout_callbacks = [] + + async def execute(self, command): + self.commands.append(command) + if isinstance(command, RunPipelineInput): + if self.run_pipeline_side_effect is not None: + await self.run_pipeline_side_effect() + return self.pipeline_result + if isinstance(command, EndServiceOnce): + return self.end_reply + if isinstance(command, ExportSession): + return None + if isinstance(command, NotifyBridgeDone): + return None + if isinstance(command, NotifyBridgeStop): + return None + if isinstance(command, SetPendingInterrupt): + return None + if command.__class__.__name__ == "StartSpeech": + if self.speech_handles: + return self.speech_handles.pop(0) + return self.speech_handle + if command.__class__.__name__ == "WaitForSpeechPlayout": + if self.wait_for_playout_callbacks: + callback = self.wait_for_playout_callbacks.pop(0) + if callback is not None: + result = callback(command) + if asyncio.iscoroutine(result): + await result + if self.wait_for_playout_exceptions: + exc = self.wait_for_playout_exceptions.pop(0) + if exc is not None: + raise exc + delay_s = ( + self.wait_for_playout_delays_s.pop(0) + if self.wait_for_playout_delays_s + else self.wait_for_playout_delay_s + ) + if delay_s > 0: + await asyncio.sleep(delay_s) + return None + if isinstance(command, InterruptSpeech): + setattr(command.handle, "interrupted", True) + return None + if command.__class__.__name__ == "SetPipelineInterruption": + return None + if command.__class__.__name__ == "InjectIdleNudge": + return None + if command.__class__.__name__ == "StartSession": + if self.on_start_session is not None: + self.on_start_session(command) + return None + raise AssertionError(f"Unsupported command in fake executor: {command!r}") + + def execute_now(self, command): + self.commands.append(command) + if isinstance(command, ExtractSpokenText): + return self.spoken_text + raise AssertionError(f"Unsupported synchronous command: {command!r}") + + +class _FakeVADThresholdController: + def __init__(self) -> None: + self.calls = [] + self.active = False + + def activate_agent_wait_timeout_retry(self, *, attempt: int, reason: str) -> bool: + if self.active: + return False + self.calls.append(("activate", attempt, reason)) + self.active = True + return True + + def restore(self, *, reason: str) -> bool: + if not self.active: + return False + self.calls.append(("restore", reason)) + self.active = False + return True + + +class _FakeVAD: + def __init__(self) -> None: + self.update_options_calls = [] + + def update_options(self, **kwargs) -> None: + self.update_options_calls.append(kwargs) + + +class _AwaitableSpeechHandle: + def __init__(self) -> None: + self.awaited = False + + async def wait_for_playout(self) -> None: + return None + + def interrupt(self, *, force: bool = False) -> None: + return None + + def __await__(self): + async def _wait(): + self.awaited = True + return self + + return _wait().__await__() + + +class SpeechServiceTests(_AsyncTestCase): + async def test_start_returns_livekit_speech_handle_without_waiting_for_playout(self) -> None: + handle = _AwaitableSpeechHandle() + session = SimpleNamespace( + say=lambda *args, **kwargs: handle, + ) + service = SpeechService(session) + + result = await service.start( + "texto", + allow_interruptions=True, + add_to_chat_ctx=True, + ) + + self.assertIs(result, handle) + self.assertFalse(handle.awaited) + + async def test_start_passes_provided_audio_to_livekit_session(self) -> None: + handle = _AwaitableSpeechHandle() + calls = [] + audio = object() + session = SimpleNamespace( + say=lambda *args, **kwargs: calls.append((args, kwargs)) or handle, + ) + service = SpeechService(session) + + result = await service.start( + "texto", + allow_interruptions=True, + add_to_chat_ctx=False, + audio=audio, + ) + + self.assertIs(result, handle) + self.assertIs(calls[0][1]["audio"], audio) + self.assertFalse(calls[0][1]["add_to_chat_ctx"]) + + +class DynamicVADThresholdControllerTests(unittest.TestCase): + def test_activate_and_restore_updates_vad_options(self) -> None: + vad = _FakeVAD() + controller = DynamicVADThresholdController( + vad, + config=DynamicVADThresholdConfig( + enabled=True, + baseline_activation_threshold=0.30, + baseline_deactivation_threshold=0.15, + retry_activation_threshold=0.20, + retry_deactivation_threshold=0.05, + ), + ) + + self.assertEqual( + controller.current_thresholds(), + { + "mode": "baseline", + "activation_threshold": 0.30, + "deactivation_threshold": 0.15, + }, + ) + self.assertTrue( + controller.activate_agent_wait_timeout_retry( + attempt=1, + reason="agent_wait_timeout_retry", + ) + ) + self.assertEqual( + controller.current_thresholds(), + { + "mode": "agent_wait_timeout_retry", + "activation_threshold": 0.20, + "deactivation_threshold": 0.05, + }, + ) + self.assertTrue(controller.restore(reason="user_final")) + self.assertEqual( + controller.current_thresholds(), + { + "mode": "baseline", + "activation_threshold": 0.30, + "deactivation_threshold": 0.15, + }, + ) + + self.assertEqual( + vad.update_options_calls, + [ + { + "activation_threshold": 0.20, + "deactivation_threshold": 0.05, + }, + { + "activation_threshold": 0.30, + "deactivation_threshold": 0.15, + }, + ], + ) + + +class InterruptPolicyTests(unittest.TestCase): + def test_allow_stage_and_backchannel(self) -> None: + policy = InterruptPolicy() + self.assertTrue(policy.allow_stage("ARGUMENTATION")) + self.assertFalse(policy.allow_stage("FORMALIZATION")) + self.assertFalse(policy.allow_stage("IDLE_NUDGE")) + self.assertTrue(policy.allow_stage("AGENT_WAIT_TIMEOUT_RETRY")) + self.assertTrue(policy.allow_stage("unknown-stage")) + self.assertTrue(policy.should_ignore_backchannel(speaking=True, text="uhum")) + self.assertFalse(policy.should_ignore_backchannel(speaking=False, text="uhum")) + + +class IdlePolicyTests(unittest.TestCase): + def test_idle_policy_decisions(self) -> None: + policy = IdlePolicy( + IdlePolicyConfig( + enabled=True, + delay_s=10.0, + join_delay_s=15.0, + close_delay_s=20.0, + max_tries=3, + end_reason="no_user_response", + ) + ) + + self.assertTrue(policy.should_arm_nudge(delay_s=10.0, nudge_text="alo")) + self.assertFalse(policy.should_arm_nudge(delay_s=0.0, nudge_text="alo")) + self.assertTrue( + policy.should_fire_nudge( + token_is_current=True, + seq_matches=True, + speaking=False, + gap_active=False, + finalized=False, + current_stage="PRESENTATION", + ) + ) + self.assertFalse( + policy.should_fire_nudge( + token_is_current=True, + seq_matches=True, + speaking=False, + gap_active=False, + finalized=False, + current_stage="DONE", + ) + ) + self.assertTrue(policy.should_arm_close_after_nudge(stage="IDLE_NUDGE", nudge_count=3)) + + def test_idle_policy_is_disabled_by_default(self) -> None: + policy = IdlePolicy( + IdlePolicyConfig( + enabled=False, + delay_s=10.0, + join_delay_s=15.0, + close_delay_s=20.0, + max_tries=3, + end_reason="no_user_response", + ) + ) + + self.assertFalse(policy.should_arm_nudge(delay_s=10.0, nudge_text="alo")) + self.assertFalse(policy.should_arm_close_before_fire(nudge_count=3)) + self.assertFalse(policy.should_arm_close_after_nudge(stage="IDLE_NUDGE", nudge_count=3)) + + +class FinalizationPolicyTests(unittest.TestCase): + def test_finalization_decisions(self) -> None: + policy = FinalizationPolicy(call_end_grace_s=2.0) + self.assertTrue(policy.should_skip(finalized=True)) + self.assertFalse(policy.should_skip(finalized=False)) + self.assertTrue(policy.should_finalize_room_empty(remote_participants=0, finalized=False)) + self.assertFalse(policy.should_finalize_room_empty(remote_participants=1, finalized=False)) + self.assertTrue(policy.is_done_stage("DONE")) + self.assertFalse(policy.is_done_stage("PRESENTATION")) + + +class TimerSchedulerTests(_AsyncTestCase): + async def test_arm_updates_token_and_cancel_stops_task(self) -> None: + created_tasks = [] + logger = _NullLogger() + + def _create_task_logged(coro, *, name: str): + task = asyncio.create_task(coro, name=name) + created_tasks.append(task) + return task + + scheduler = TimerScheduler(create_task_logged=_create_task_logged, logger=logger) + gate = asyncio.Event() + + async def _wait_forever(): + await gate.wait() + + token1 = scheduler.arm("idle", task_name="idle_1", coro=_wait_forever()) + token2 = scheduler.arm("idle", task_name="idle_2", coro=_wait_forever()) + + self.assertFalse(scheduler.is_current("idle", token1)) + self.assertTrue(scheduler.is_current("idle", token2)) + + scheduler.cancel("idle", reason="test", log_name="IDLE_TIMER_CANCEL") + await asyncio.sleep(0) + + self.assertTrue(created_tasks[-1].cancelled()) + gate.set() + await asyncio.gather(*created_tasks, return_exceptions=True) + + +class RuntimeCommandExecutorTests(_AsyncTestCase): + async def test_routes_commands_to_underlying_services(self) -> None: + agent = _FakeAgent() + bridge_gateway = SimpleNamespace(notify_stage_done=self._make_async_noop()) + bridge_gateway.notify_stop = self._make_async_noop() + export_service = SimpleNamespace(export_session=self._make_async_noop()) + session = _FakeSession() + speech_service = SimpleNamespace( + start=self._make_async_return("speech-handle"), + wait_for_playout=self._make_async_noop(), + extract_spoken_text=lambda handle: "spoken", + ) + agent.pipeline = SimpleNamespace( + inject_idle_nudge=self._make_async_noop(), + set_interruption=self._make_async_noop(), + run=self._make_async_return(BackendReply(stage="DONE", text="encerrando", done=True)), + ) + + executor = RuntimeCommandExecutor( + agent=agent, + bridge_gateway=bridge_gateway, + export_service=export_service, + session=session, + speech_service=speech_service, + ) + + result = await executor.execute(RunPipelineInput("payload")) + self.assertEqual(result, BackendReply(stage="DONE", text="encerrando", done=True)) + self.assertEqual(executor.execute_now(ExtractSpokenText(object())), "spoken") + + @staticmethod + def _make_async_noop(): + async def _noop(*args, **kwargs): + return None + + return _noop + + @staticmethod + def _make_async_return(value): + async def _return(*args, **kwargs): + return value + + return _return + + +class CallRuntimeTests(_AsyncTestCase): + def setUp(self) -> None: + self.created_tasks = [] + + def _create_task_logged(self, coro, *, name: str): + task = asyncio.create_task(coro, name=name) + self.created_tasks.append(task) + return task + + async def _drain_tasks(self) -> None: + while True: + pending = [task for task in self.created_tasks if not task.done()] + if not pending: + return + await asyncio.gather(*pending) + + async def _cancel_pending_tasks(self) -> None: + pending = [task for task in self.created_tasks if not task.done()] + for task in pending: + task.cancel() + if pending: + await asyncio.gather(*pending, return_exceptions=True) + + def test_inflight_backend_wait_message_id_uses_uuid_when_original_is_missing(self) -> None: + generated = uuid.UUID("12345678-1234-4234-9234-123456789abc") + + with mock.patch("app.livekit.runtime.call_runtime.uuid.uuid4", return_value=generated): + message_id = CallRuntime._inflight_backend_wait_message_id(attempt=1) + + self.assertEqual( + message_id, + "12345678-1234-4234-9234-123456789abc_conforto_12345678123442349234123456789abc", + ) + + def test_feedback_message_id_keeps_the_agent_speech_id(self) -> None: + message_id = CallRuntime._feedback_message_id( + message_id="turno-1", + speech_id="speech-feedback-1", + ) + + self.assertEqual(message_id, "turno-1_feedback_speech-feedback-1") + self.assertEqual(CallRuntime._base_message_id(message_id), "turno-1") + + def test_auxiliary_message_ids_are_distinct_and_keep_the_base(self) -> None: + generated = uuid.UUID("12345678-1234-4234-9234-123456789abc") + with mock.patch("app.livekit.runtime.call_runtime.uuid.uuid4", return_value=generated): + interruption = CallRuntime._interruption_comfort_message_id(message_id="turno-1") + tts_error = CallRuntime._tts_error_message_id(message_id="turno-1") + + self.assertEqual(CallRuntime._idle_nudge_message_id(message_id="turno-1", attempt=2), "turno-1_idle_nudge_2") + self.assertEqual(interruption, "turno-1_interruption_confort_12345678123442349234123456789abc") + self.assertEqual(tts_error, "turno-1_tts_error_12345678123442349234123456789abc") + self.assertEqual(CallRuntime._base_message_id(interruption), "turno-1") + self.assertEqual(CallRuntime._base_message_id(tts_error), "turno-1") + self.assertTrue( + CallRuntime._transfer_message_id(message_id="turno-1").startswith( + "turno-1_transferencia_" + ) + ) + self.assertTrue( + CallRuntime._resource_error_message_id( + message_id="turno-1", + terminal=True, + tipo_evento="envio msg", + ).startswith("turno-1_erro_terminal_envio_") + ) + + async def test_backend_feedback_uses_a_suffixed_message_id(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + + with mock.patch.object(runtime, "say_stage", new_callable=mock.AsyncMock) as say_stage: + await runtime._speak_backend_reply( + BackendReply( + stage="PRESENTATION", + text="Ainda estou consultando.", + metadata={ + "message_id": "turno-feedback", + "speech_id": "speech-feedback-1", + "agent_message_type": "feedback", + }, + ), + source="backend_push", + add_to_chat_ctx=True, + ) + + self.assertEqual( + say_stage.await_args.kwargs["message_id"], + "turno-feedback_feedback_speech-feedback-1", + ) + self.assertEqual(say_stage.await_args.kwargs["speech_id"], "speech-feedback-1") + + def _make_runtime( + self, + *, + agent_starts_conversation: bool = False, + push_replies=None, + runtime_config_overrides=None, + idle_policy_overrides=None, + vad_threshold_controller=None, + ): + agent = _FakeAgent() + if push_replies is not None: + agent.pipeline = _FakePushPipeline(push_replies) + session = _FakeSession() + ctx = _FakeContext() + executor = _FakeCommandExecutor() + config_kwargs = { + "call_end_grace_s": 0.0, + "final_grace_s": 0.0, + "idle_nudge_delay_s": 0.0, + "idle_nudge_join_delay_s": 0.0, + "idle_nudge_close_delay_s": 0.0, + "idle_nudge_max_tries": 0, + "idle_nudge_end_reason": "no_user_response", + } + if runtime_config_overrides: + config_kwargs.update(runtime_config_overrides) + idle_policy_kwargs = { + "enabled": False, + "delay_s": 0.0, + "join_delay_s": 0.0, + "close_delay_s": 0.0, + "max_tries": 0, + "end_reason": "no_user_response", + } + if idle_policy_overrides: + idle_policy_kwargs.update(idle_policy_overrides) + runtime = CallRuntime( + ctx=ctx, + session=session, + agent=agent, + command_executor=executor, + call_logger=_NullLogger(), + protocol="PRT-1", + session_id="S-1", + bridge_identity="bridge-1", + agent_starts_conversation=agent_starts_conversation, + nudge_text="alo", + call_t0=0.0, + extract_text_from_transcript=lambda transcript: transcript, + interrupt_policy=InterruptPolicy(), + idle_policy=IdlePolicy(IdlePolicyConfig(**idle_policy_kwargs)), + finalization_policy=FinalizationPolicy(call_end_grace_s=0.0), + create_task_logged=self._create_task_logged, + config=RuntimeConfig(**config_kwargs), + logger=_NullLogger(), + vad_threshold_controller=vad_threshold_controller, + ) + return runtime, agent, executor + + def test_inflight_backend_wait_uses_short_audio_then_long_audio(self) -> None: + with TemporaryDirectory() as tmp_dir: + base_dir = Path(tmp_dir) + short_dir = base_dir / "short" + long_dir = base_dir / "long" + short_dir.mkdir() + long_dir.mkdir() + short_audio = short_dir / "curto.wav" + long_audio = long_dir / "longo.wav" + short_audio.write_bytes(b"short") + long_audio.write_bytes(b"long") + + runtime, _agent, _executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_short_audio_dir": str(short_dir), + "inflight_backend_wait_long_audio_dir": str(long_dir), + } + ) + + self.assertEqual( + runtime._inflight_backend_wait_audio_path(attempt=1), + short_audio, + ) + self.assertEqual( + runtime._inflight_backend_wait_audio_path(attempt=2), + long_audio, + ) + + def test_inflight_backend_wait_prewarms_audio_duration_cache(self) -> None: + with TemporaryDirectory() as tmp_dir: + short_dir = Path(tmp_dir) / "short" + short_dir.mkdir() + short_audio = short_dir / "curto.wav" + short_audio.write_bytes(b"fake-wav") + + with mock.patch( + "app.livekit.runtime.call_runtime.wav_duration_ms", + return_value=321, + ) as duration_ms: + runtime, _agent, _executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_short_audio_dir": str(short_dir), + } + ) + + duration_ms.assert_called_once_with(str(short_audio)) + with mock.patch( + "app.livekit.runtime.call_runtime.wav_duration_ms", + side_effect=AssertionError("duration should come from cache"), + ): + self.assertEqual( + runtime._inflight_backend_wait_audio_duration_ms(short_audio), + 321, + ) + + def test_inflight_backend_wait_keeps_legacy_long_audio_path_fallback(self) -> None: + runtime, _agent, _executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + + self.assertEqual( + runtime._inflight_backend_wait_audio_path(), + Path(WAIT_LONG_AUDIO_PATH), + ) + + async def test_idle_nudge_does_not_activate_retry_vad_threshold(self) -> None: + controller = _FakeVADThresholdController() + runtime, _agent, executor = self._make_runtime( + idle_policy_overrides={ + "enabled": True, + "delay_s": 10.0, + "join_delay_s": 10.0, + "close_delay_s": 10.0, + "max_tries": 2, + "end_reason": "no_user_response", + }, + vad_threshold_controller=controller, + ) + + try: + runtime.arm_idle_timer(reason="test", delay_s=0.001) + for _ in range(50): + if any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands): + break + await asyncio.sleep(0.001) + + self.assertNotIn(("activate", 1, "idle_nudge_fire"), controller.calls) + start_speech = next( + cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech" + ) + self.assertEqual(start_speech.text, "alo") + self.assertFalse(start_speech.allow_interruptions) + + runtime.on_user_input_transcribed(SimpleNamespace(transcript="sim", is_final=True)) + self.assertNotIn(("restore", "user_final"), controller.calls) + finally: + await self._cancel_pending_tasks() + + async def test_idle_nudge_uses_a_suffixed_message_id(self) -> None: + runtime, _agent, _executor = self._make_runtime( + idle_policy_overrides={ + "enabled": True, + "delay_s": 10.0, + "join_delay_s": 10.0, + "close_delay_s": 10.0, + "max_tries": 2, + "end_reason": "no_user_response", + } + ) + runtime._last_agent_message_id = "turno-idle" + + try: + with mock.patch.object(runtime, "say_stage", new_callable=mock.AsyncMock) as say_stage: + runtime.arm_idle_timer(reason="test", delay_s=0.001) + for _ in range(50): + if say_stage.await_count: + break + await asyncio.sleep(0.001) + + say_stage.assert_awaited_once() + self.assertEqual( + say_stage.await_args.kwargs["message_id"], + "turno-idle_idle_nudge_1", + ) + finally: + await self._cancel_pending_tasks() + + async def test_on_user_input_transcribed_stashes_when_stage_is_not_interruptible(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.current_stage = "FORMALIZATION" + runtime.speaking.set() + + event = SimpleNamespace(transcript="quero continuar", is_final=True) + with mock.patch("app.livekit.runtime.call_runtime.log_flow_event") as flow_log: + runtime.on_user_input_transcribed(event) + await self._drain_tasks() + + self.assertIsNotNone(runtime.state.pending_user_final) + self.assertEqual(runtime.state.pending_user_final.text, "quero continuar") + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertTrue( + any( + call.args[1] == "interrupt_ignored" + and call.kwargs["reason"] == "stage_not_interruptible" + for call in flow_log.call_args_list + ) + ) + + async def test_on_user_input_transcribed_runs_when_argumentation_is_speaking(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.current_stage = "ARGUMENTATION" + runtime.state.current_speech.allow_interruptions = True + runtime.speaking.set() + + event = SimpleNamespace(transcript="quero continuar", is_final=True) + runtime.on_user_input_transcribed(event) + await self._drain_tasks() + + self.assertIsNone(runtime.state.pending_user_final) + self.assertTrue(any(isinstance(cmd, SetPendingInterrupt) for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_short_vad_utterance_with_text_is_discarded_while_r1_is_processing(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + runtime.state.deferred_interruption.backend_in_flight = True + register_transcribed_turn( + runtime._structured_log_context, + message_id="short-utterance", + transcription="sim", + text="sim", + ) + runtime.note_vad_speech_end(999) + runtime.on_user_input_transcribed(SimpleNamespace(transcript="sim", is_final=True)) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 1) + self.assertEqual(runtime.state.deferred_interruption.long_turns, []) + self.assertFalse(runtime.state.deferred_interruption.special_comfort_sent) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_first_long_vad_plays_special_before_empty_stt(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + runtime.state.deferred_interruption.backend_in_flight = True + with mock.patch.object(runtime, "_play_deferred_interruption_comfort") as comfort: + runtime.note_vad_speech_end(1000) + comfort.assert_called_once_with() + runtime.on_user_input_transcribed(SimpleNamespace(transcript="", is_final=True)) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 1) + self.assertEqual(runtime.state.deferred_interruption.long_turns, []) + self.assertTrue(runtime.state.deferred_interruption.special_comfort_sent) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_special_processing_comfort_starts_protected(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.deferred_interruption.special_comfort_sent = True + + runtime._play_deferred_interruption_comfort() + await self._drain_tasks() + + speech = next(cmd for cmd in executor.commands if isinstance(cmd, StartSpeech)) + self.assertEqual(speech.text, "Um instante") + self.assertFalse(speech.allow_interruptions) + + async def test_short_backend_wait_comfort_starts_protected(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + + with mock.patch.object( + runtime, + "_select_inflight_backend_wait_audio", + return_value=(Path(WAIT_LONG_AUDIO_PATH), "short"), + ): + played = await runtime._play_inflight_backend_wait_notice_audio( + user_seq=None, + attempt=1, + ) + + self.assertTrue(played) + speech = next(cmd for cmd in executor.commands if isinstance(cmd, StartSpeech)) + self.assertFalse(speech.allow_interruptions) + + async def test_short_comfort_is_not_rearmed_by_its_originating_user_final(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.speaking.set() + runtime.state.current_speech.stage = "AGENT_BACKEND_WAIT" + runtime.state.current_speech.handle = executor.speech_handle + runtime.state.current_speech.allow_interruptions = False + runtime.state.current_speech.rearm_interruptions_on_next_user_speech = True + executor.speech_handle.allow_interruptions = False + + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="fala original", is_final=True) + ) + await self._drain_tasks() + + self.assertFalse(runtime.state.current_speech.allow_interruptions) + self.assertFalse(executor.speech_handle.allow_interruptions) + self.assertTrue( + runtime.state.current_speech.rearm_interruptions_on_next_user_speech + ) + + async def test_short_comfort_is_rearmed_on_new_user_speech(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.speaking.set() + runtime.state.current_speech.stage = "INTERRUPTION_COMFORT" + runtime.state.current_speech.handle = executor.speech_handle + runtime.state.current_speech.allow_interruptions = False + runtime.state.current_speech.rearm_interruptions_on_next_user_speech = True + executor.speech_handle.allow_interruptions = False + + runtime._rearm_current_speech_interruptions_on_new_user_speech() + + self.assertTrue(runtime.state.current_speech.allow_interruptions) + self.assertFalse( + runtime.state.current_speech.rearm_interruptions_on_next_user_speech + ) + self.assertTrue(executor.speech_handle.allow_interruptions) + + async def test_feedback_compat_mode_drops_processing_speech_without_replay(self) -> None: + runtime, _agent, executor = self._make_runtime( + runtime_config_overrides={"deferred_interruption_enabled": False} + ) + runtime.state.user_final_seq = 1 + runtime.state.deferred_interruption.backend_in_flight = True + register_transcribed_turn( + runtime._structured_log_context, + message_id="processing-drop", + transcription="quero interromper", + text="quero interromper", + ) + + with ( + mock.patch.object(runtime, "_play_deferred_interruption_comfort") as special, + mock.patch.object(runtime, "_play_deferred_short_comfort") as short, + mock.patch("app.livekit.runtime.call_runtime.log_flow_event") as flow_log, + ): + runtime.note_vad_speech_end(2500) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="quero interromper", is_final=True) + ) + await self._drain_tasks() + + special.assert_not_called() + short.assert_not_called() + self.assertEqual(runtime.state.user_final_seq, 1) + self.assertEqual(runtime.state.deferred_interruption.long_turns, []) + self.assertEqual(len(runtime.state.deferred_interruption.pending_vad_utterances), 0) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertTrue( + any( + call.args[1] == "backend_processing_vad_discarded" + and call.kwargs["mode"] == "feedback_compat" + for call in flow_log.call_args_list + ) + ) + self.assertTrue( + any( + call.args[1] == "user_input_dropped" + and call.kwargs["reason"] == "backend_processing_feedback_compat" + and call.kwargs["agent_message_type"] == "feedback" + for call in flow_log.call_args_list + ) + ) + + async def test_vad_end_outside_backend_window_is_not_deferred(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime.state.deferred_interruption.backend_in_flight = False + runtime.speaking.set() + + with mock.patch.object(runtime, "_play_deferred_interruption_comfort") as comfort: + runtime.note_vad_speech_end(4000) + + comfort.assert_not_called() + self.assertEqual(len(runtime.state.deferred_interruption.pending_vad_utterances), 0) + self.assertFalse(runtime.state.deferred_interruption.special_comfort_sent) + + async def test_barge_in_during_agent_playout_still_runs_a_new_pipeline(self) -> None: + runtime, agent, executor = self._make_runtime() + runtime.state.current_stage = "ARGUMENTATION" + runtime.state.current_speech.allow_interruptions = True + runtime.speaking.set() + # O playout segura o run_lock, mas a janela do ciclo diferido ja fechou: + # a fala precisa voltar a ser barge-in comum. + await agent._run_lock.acquire() + try: + runtime.note_vad_speech_end(4000) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="muda de assunto", is_final=True) + ) + await asyncio.sleep(0) + finally: + agent._run_lock.release() + await self._drain_tasks() + + self.assertEqual(runtime.state.deferred_interruption.long_turns, []) + self.assertTrue(any(isinstance(cmd, SetPendingInterrupt) for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_long_vad_transcripts_are_concatenated_once_and_r2_blocks_replay(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + runtime.state.deferred_interruption.backend_in_flight = True + with ( + mock.patch.object(runtime, "_play_deferred_interruption_comfort") as comfort, + mock.patch.object(runtime, "_play_deferred_short_comfort") as short_comfort, + mock.patch.object(runtime, "run_pipeline", new_callable=mock.AsyncMock) as run_pipeline, + ): + runtime.note_vad_speech_end(1200) + runtime.note_vad_speech_end(300) + runtime.note_vad_speech_end(1400) + register_transcribed_turn( + runtime._structured_log_context, + message_id="long-1", + transcription="primeira fala longa", + text="primeira fala longa", + ) + register_transcribed_turn( + runtime._structured_log_context, + message_id="short-1", + transcription="fala curta", + text="fala curta", + ) + register_transcribed_turn( + runtime._structured_log_context, + message_id="long-2", + transcription="segunda fala longa", + text="segunda fala longa", + ) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="primeira fala longa", is_final=True) + ) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="fala curta", is_final=True) + ) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="segunda fala longa", is_final=True) + ) + + comfort.assert_called_once_with() + short_comfort.assert_called_once_with(source="repeated_interruption") + self.assertEqual( + [turn.text for turn in runtime.state.deferred_interruption.long_turns], + ["primeira fala longa", "segunda fala longa"], + ) + dispatched = await runtime._dispatch_deferred_interruption( + BackendReply(stage="PRESENTATION", text="resposta R1") + ) + await self._drain_tasks() + + self.assertTrue(dispatched) + self.assertTrue(runtime.state.deferred_interruption.replay_in_flight) + run_pipeline.assert_awaited_once_with( + "primeira fala longa. segunda fala longa", + "primeira fala longa. segunda fala longa", + is_deferred_replay=True, + user_seq=3, + message_id="long-2", + inflight_initial_notice_consumed=True, + ) + + runtime.note_vad_speech_end(1500) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="nao deve gerar R3", is_final=True) + ) + self.assertEqual(runtime.state.user_final_seq, 3) + self.assertEqual(runtime.state.deferred_interruption.long_turns, []) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + + async def test_vad_ended_before_r1_is_promoted_when_stt_finishes_during_r1(self) -> None: + runtime, _agent, executor = self._make_runtime() + backend_started = asyncio.Event() + release_backend = asyncio.Event() + runtime.state.user_final_seq = 0 + + async def _wait_for_backend() -> None: + backend_started.set() + await release_backend.wait() + + executor.run_pipeline_side_effect = _wait_for_backend + with ( + mock.patch.object(runtime, "_play_deferred_interruption_comfort"), + mock.patch.object( + runtime, + "_dispatch_deferred_interruption", + new_callable=mock.AsyncMock, + return_value=True, + ) as dispatch, + ): + # O segundo fim de fala acontece antes de R1, mas seu STT chega + # somente enquanto R1 está em voo. + runtime.note_vad_speech_end(1500) + runtime.note_vad_speech_end(1500) + register_transcribed_turn( + runtime._structured_log_context, + message_id="r1-turn", + transcription="primeiro turno", + text="primeiro turno", + ) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="primeiro turno", is_final=True) + ) + await backend_started.wait() + + register_transcribed_turn( + runtime._structured_log_context, + message_id="late-turn", + transcription="fala que terminou antes de R1", + text="fala que terminou antes de R1", + ) + runtime.on_user_input_transcribed( + SimpleNamespace( + transcript="fala que terminou antes de R1", is_final=True + ) + ) + release_backend.set() + # O task criado pelo callback pode precisar de um ciclo extra para + # concluir a espera e tentar o dispatch. + await self._drain_tasks() + + dispatch.assert_awaited_once() + self.assertEqual(len(runtime.state.deferred_interruption.pre_backend_vad_utterances), 0) + + async def test_r1_waits_for_pending_long_stt_before_dispatching(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + backend_started = asyncio.Event() + release_backend = asyncio.Event() + + async def _wait_for_backend() -> None: + backend_started.set() + await release_backend.wait() + + executor.run_pipeline_side_effect = _wait_for_backend + with ( + mock.patch.object(runtime, "_play_deferred_interruption_comfort"), + mock.patch.object(runtime, "_dispatch_deferred_interruption", new_callable=mock.AsyncMock, return_value=True) as dispatch, + ): + task = asyncio.create_task(runtime.run_pipeline("original", "original", user_seq=1)) + await backend_started.wait() + runtime.note_vad_speech_end(1200) + release_backend.set() + await asyncio.sleep(0) + + self.assertFalse(task.done()) + self.assertFalse(any(isinstance(cmd, StartSpeech) for cmd in executor.commands)) + register_transcribed_turn( + runtime._structured_log_context, + message_id="long-wait", + transcription="fala longa", + text="fala longa", + ) + runtime.on_user_input_transcribed(SimpleNamespace(transcript="fala longa", is_final=True)) + await task + + dispatch.assert_awaited_once() + + async def test_r2_discards_input_until_its_playout_finishes(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime.state.deferred_interruption.replay_in_flight = True + executor.pipeline_result = BackendReply(stage="PRESENTATION", text="resposta R2") + played = [] + + def _during_r2_playout(_command) -> None: + played.append(True) + self.assertTrue(runtime.state.deferred_interruption.replay_in_flight) + runtime.note_vad_speech_end(1500) + runtime.on_user_input_transcribed(SimpleNamespace(transcript="fala ignorada", is_final=True)) + + executor.wait_for_playout_callbacks.append(_during_r2_playout) + with mock.patch.object(runtime, "_play_deferred_short_comfort") as short_comfort: + await runtime.run_pipeline("texto R2", "texto R2", user_seq=1, is_deferred_replay=True) + + self.assertTrue(played, "o playout de R2 precisa acontecer de verdade") + short_comfort.assert_not_called() + self.assertFalse(runtime.state.deferred_interruption.replay_in_flight) + self.assertEqual(runtime.state.user_final_seq, 1) + self.assertEqual(runtime.state.deferred_interruption.long_turns, []) + + async def test_r2_reply_is_not_interruptible(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.current_stage = "ARGUMENTATION" + runtime.state.user_final_seq = 1 + runtime.state.deferred_interruption.replay_in_flight = True + executor.pipeline_result = BackendReply(stage="ARGUMENTATION", text="resposta R2") + + await runtime.run_pipeline("texto R2", "texto R2", user_seq=1, is_deferred_replay=True) + + speech_commands = [cmd for cmd in executor.commands if isinstance(cmd, StartSpeech)] + self.assertEqual([cmd.text for cmd in speech_commands], ["resposta R2"]) + # Descartar a fala no runtime nao basta: com allow_interruptions=True o + # barge-in do proprio LiveKit corta o playout e R2 sai pela metade. + self.assertFalse(speech_commands[0].allow_interruptions) + + async def test_r2_logs_and_ignores_short_voice_without_playing_comfort(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime.state.deferred_interruption.replay_in_flight = True + + with ( + mock.patch.object(runtime, "_play_deferred_short_comfort") as short_comfort, + mock.patch("app.livekit.runtime.call_runtime.log_flow_event") as flow_log, + ): + runtime.note_vad_speech_end(999) + + short_comfort.assert_not_called() + self.assertTrue( + any( + call.args[1] == "deferred_replay_vad_discarded" + and call.kwargs["speech_duration_ms"] == 999 + and not call.kwargs["short_comfort_scheduled"] + for call in flow_log.call_args_list + ) + ) + + async def test_normal_reply_stays_interruptible_outside_the_replay(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.current_stage = "ARGUMENTATION" + runtime.state.user_final_seq = 1 + executor.pipeline_result = BackendReply(stage="ARGUMENTATION", text="resposta normal") + + await runtime.run_pipeline("texto", "texto", user_seq=1) + + speech_commands = [cmd for cmd in executor.commands if isinstance(cmd, StartSpeech)] + self.assertEqual([cmd.text for cmd in speech_commands], ["resposta normal"]) + self.assertTrue(speech_commands[0].allow_interruptions) + + async def test_r1_waits_for_the_user_turn_to_end_before_deciding(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + backend_started = asyncio.Event() + release_backend = asyncio.Event() + + async def _wait_for_backend() -> None: + backend_started.set() + await release_backend.wait() + + executor.run_pipeline_side_effect = _wait_for_backend + with ( + mock.patch.object(runtime, "_play_deferred_interruption_comfort"), + mock.patch.object( + runtime, + "_dispatch_deferred_interruption", + new_callable=mock.AsyncMock, + return_value=True, + ) as dispatch, + ): + task = asyncio.create_task(runtime.run_pipeline("original", "original", user_seq=1)) + await backend_started.wait() + # R1 volta com o cliente no meio da fala: nem o VAD fechou, nem ha + # STT pendente ainda. + runtime._user_not_speaking.clear() + release_backend.set() + await asyncio.sleep(0) + + self.assertFalse(task.done()) + self.assertFalse(any(isinstance(cmd, StartSpeech) for cmd in executor.commands)) + + runtime.note_vad_speech_end(1300) + runtime._user_not_speaking.set() + register_transcribed_turn( + runtime._structured_log_context, + message_id="late-long", + transcription="fala que comecou antes de R1 voltar", + text="fala que comecou antes de R1 voltar", + ) + runtime.on_user_input_transcribed( + SimpleNamespace( + transcript="fala que comecou antes de R1 voltar", is_final=True + ) + ) + await task + + dispatch.assert_awaited_once() + self.assertFalse(any(isinstance(cmd, StartSpeech) for cmd in executor.commands)) + + async def test_r1_gives_up_waiting_when_the_stt_final_never_arrives(self) -> None: + runtime, _agent, executor = self._make_runtime( + runtime_config_overrides={ + "deferred_interruption_stt_settle_timeout_s": 0.01, + } + ) + runtime.state.user_final_seq = 1 + backend_started = asyncio.Event() + release_backend = asyncio.Event() + + async def _wait_for_backend() -> None: + backend_started.set() + await release_backend.wait() + + executor.run_pipeline_side_effect = _wait_for_backend + executor.pipeline_result = BackendReply(stage="PRESENTATION", text="resposta R1") + with ( + mock.patch.object(runtime, "_play_deferred_interruption_comfort"), + mock.patch.object( + runtime, "_speak_backend_reply", new_callable=mock.AsyncMock + ) as speak, + ): + task = asyncio.create_task(runtime.run_pipeline("original", "original", user_seq=1)) + await backend_started.wait() + runtime.note_vad_speech_end(1200) + release_backend.set() + # Sem teto na espera esta chamada nunca retornaria. + await asyncio.wait_for(task, timeout=5) + + deferred = runtime.state.deferred_interruption + self.assertEqual(deferred.pending_long_stt_finals, 0) + self.assertEqual(len(deferred.pending_vad_utterances), 0) + self.assertFalse(deferred.backend_in_flight) + # A resposta de R1 e liberada em vez de ficar presa a um final que nao vem. + speak.assert_awaited_once() + + async def test_special_comfort_flag_is_cleared_when_r1_answers_empty(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime.state.user_final_seq = 1 + backend_started = asyncio.Event() + release_backend = asyncio.Event() + + async def _wait_for_backend() -> None: + backend_started.set() + await release_backend.wait() + + _executor.run_pipeline_side_effect = _wait_for_backend + _executor.pipeline_result = BackendReply(stage="PRESENTATION", text="") + with mock.patch.object(runtime, "_play_deferred_interruption_comfort"): + task = asyncio.create_task(runtime.run_pipeline("original", "original", user_seq=1)) + await backend_started.wait() + runtime.note_vad_speech_end(1200) + runtime.on_user_input_transcribed(SimpleNamespace(transcript="", is_final=True)) + release_backend.set() + await task + + # Sem a limpeza aqui o conforto especial nunca mais tocaria na ligacao. + self.assertFalse(runtime.state.deferred_interruption.special_comfort_sent) + + async def test_empty_final_is_counted_once_when_provider_and_livekit_both_report(self) -> None: + runtime, _agent, _executor = self._make_runtime() + # Em raw_json o transcript cru nao e vazio, so o texto extraido dele e. + runtime._extract_text_from_transcript = lambda _transcript: "" + deferred = runtime.state.deferred_interruption + deferred.backend_in_flight = True + runtime.note_vad_speech_end(300) + runtime.note_vad_speech_end(1500) + + with mock.patch.object(runtime, "_play_deferred_interruption_comfort"): + # Provider avisa o final vazio da fala curta... + runtime.handle_empty_stt_final(source="stt_provider") + # ...e o LiveKit repassa o mesmo resultado como evento. + runtime.on_user_input_transcribed( + SimpleNamespace(transcript='{"data":{"text":""}}', is_final=True) + ) + await self._drain_tasks() + + # A fala longa continua na fila, aguardando o proprio final. + self.assertEqual(len(deferred.pending_vad_utterances), 1) + self.assertEqual(deferred.pending_long_stt_finals, 1) + + async def test_vad_utterance_without_final_does_not_leak_into_the_next_turn(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + backend_started = asyncio.Event() + release_backend = asyncio.Event() + + async def _wait_for_backend() -> None: + backend_started.set() + await release_backend.wait() + + executor.run_pipeline_side_effect = _wait_for_backend + with ( + mock.patch.object(runtime, "_play_deferred_interruption_comfort"), + mock.patch.object(runtime, "_speak_backend_reply", new_callable=mock.AsyncMock), + ): + task = asyncio.create_task(runtime.run_pipeline("original", "original", user_seq=1)) + await backend_started.wait() + # Ruido curto que o STT nunca finaliza. + runtime.note_vad_speech_end(200) + release_backend.set() + await asyncio.wait_for(task, timeout=5) + + deferred = runtime.state.deferred_interruption + self.assertEqual(len(deferred.pending_vad_utterances), 0) + + # No turno seguinte a fala do cliente nao pode ser consumida pela sobra. + deferred.backend_in_flight = True + with mock.patch.object(runtime, "_play_deferred_interruption_comfort"): + runtime.note_vad_speech_end(1500) + register_transcribed_turn( + runtime._structured_log_context, + message_id="turno-seguinte", + transcription="agora eu quero falar", + text="agora eu quero falar", + ) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="agora eu quero falar", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual( + [turn.text for turn in deferred.long_turns], ["agora eu quero falar"] + ) + + async def test_protected_speech_drop_wins_over_the_deferred_cycle(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime.state.user_final_seq = 1 + deferred = runtime.state.deferred_interruption + deferred.backend_in_flight = True + runtime._mark_drop_next_user_final( + {"reason": "protected_speech", "stage": "PRESENTATION"} + ) + + with mock.patch.object(runtime, "_play_deferred_interruption_comfort"): + runtime.note_vad_speech_end(1500) + register_transcribed_turn( + runtime._structured_log_context, + message_id="protegida", + transcription="fala descartada", + text="fala descartada", + ) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="fala descartada", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual(deferred.long_turns, []) + self.assertEqual(runtime.state.user_final_seq, 1) + # A fala VAD tem de sair da fila junto, senao o proximo final consumiria + # a entrada errada. + self.assertEqual(len(deferred.pending_vad_utterances), 0) + self.assertEqual(deferred.pending_long_stt_finals, 0) + + async def test_special_comfort_defers_the_periodic_backend_wait_notice(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime.state.deferred_interruption.backend_in_flight = True + runtime._inflight_backend_activity_at = 0.0 + + with mock.patch.object(runtime, "_play_deferred_interruption_comfort"): + runtime.note_vad_speech_end(1500) + + self.assertGreater(runtime._inflight_backend_activity_at, 0.0) + + def test_combined_interruption_transcription_concatenates_text_and_drops_word_timestamps(self) -> None: + turns = [ + SimpleNamespace( + transcription='{"data":{"text":"primeira", "words":[{"word":"primeira"}]}}', + text="primeira", + ), + SimpleNamespace(transcription="segunda", text="segunda"), + ] + + transcription, text = CallRuntime._combined_interruption_transcription(turns) + payload = __import__("json").loads(transcription) + + self.assertEqual(text, "primeira. segunda") + self.assertEqual(payload["data"]["text"], "primeira. segunda") + self.assertNotIn("words", payload["data"]) + + async def test_on_user_input_transcribed_drops_during_protected_feedback(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + runtime.state.current_stage = "PRESENTATION" + runtime.state.current_speech.stage = "PRESENTATION" + runtime.state.current_speech.drop_user_input_while_speaking = True + runtime.state.current_speech.agent_message_type = "feedback" + runtime.speaking.set() + + with mock.patch("app.livekit.runtime.call_runtime.log_flow_event") as flow_log: + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="quero falar agora", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 0) + self.assertIsNone(runtime.state.pending_user_final) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertTrue( + any( + call.args[1] == "user_input_dropped" + and call.kwargs["agent_message_type"] == "feedback" + and call.kwargs["text"] == "quero falar agora" + for call in flow_log.call_args_list + ) + ) + self.assertTrue( + any(event == "user_input_dropped" for event, _fields in runtime._timeline.events) + ) + + async def test_dropped_user_final_does_not_leak_message_id_to_repeated_text(self) -> None: + runtime, _agent, executor = self._make_runtime() + structured_context = StructuredLogContext( + callid="call-1", + session_id="session-drop-repeated-text", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + runtime._structured_log_context = structured_context + reset_turn_message_sequence(structured_context, clear_pending=True) + + register_started_turn_message_id(structured_context, "MSG-drop-0001") + register_transcribed_turn( + structured_context, + message_id="MSG-drop-0001", + transcription="mesmo texto", + text="mesmo texto", + ) + runtime._drop_next_user_final_context = { + "reason": "test_drop", + "stage": "PRESENTATION", + "agent_message_type": "ready", + "agent_result_type": "", + } + + try: + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="mesmo texto", is_final=True) + ) + await self._drain_tasks() + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + register_started_turn_message_id(structured_context, "MSG-current-0002") + register_transcribed_turn( + structured_context, + message_id="MSG-current-0002", + transcription="mesmo texto", + text="mesmo texto", + ) + + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="mesmo texto", is_final=True) + ) + await self._drain_tasks() + + pipeline_command = next(cmd for cmd in executor.commands if isinstance(cmd, RunPipelineInput)) + self.assertEqual(pipeline_command.user_input["message_id"], "MSG-current-0002") + finally: + reset_turn_message_sequence(structured_context, clear_pending=True) + + async def test_on_user_input_transcribed_drops_after_feedback_while_backend_is_processing(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + + await runtime._speak_backend_reply( + BackendReply( + stage="PRESENTATION", + text="Ainda estou consultando sua fatura.", + metadata={ + "event": "feedback", + "agent_message_type": "feedback", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + }, + ), + source="backend_push", + add_to_chat_ctx=True, + ) + + self.assertFalse(runtime.speaking.is_set()) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="posso falar?", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 0) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertTrue( + any( + event == "user_input_dropped" + and fields["agent_message_type"] == "feedback" + for event, fields in runtime._timeline.events + ) + ) + + async def test_on_user_input_transcribed_drops_after_result_feedback_while_backend_is_processing(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + + await runtime._speak_backend_reply( + BackendReply( + stage="PRESENTATION", + text="Aguarde um instante, por favor.", + metadata={ + "agent_message_type": "result", + "agent_result_type": "feedback", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + }, + ), + source="run_pipeline", + add_to_chat_ctx=True, + ) + + self.assertFalse(runtime.speaking.is_set()) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="voce esta cancelando?", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 0) + self.assertIsNone(runtime.state.pending_user_final) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertTrue( + any( + event == "user_input_dropped" + and fields["agent_message_type"] == "result" + and fields["agent_result_type"] == "feedback" + for event, fields in runtime._timeline.events + ) + ) + + async def test_on_user_input_transcribed_drops_during_terminal_speech(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.current_stage = "DONE" + runtime.state.current_speech.stage = "DONE" + runtime.state.current_speech.drop_user_input_while_speaking = True + runtime.state.current_speech.agent_message_type = "result" + runtime.state.current_speech.agent_result_type = "resolvido" + runtime.speaking.set() + + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="tenho outra pergunta", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 0) + self.assertIsNone(runtime.state.pending_user_final) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_user_state_speaking_during_ready_drops_late_final_but_arms_timeout(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + executor.pipeline_result = BackendReply( + stage="PRESENTATION", + text="Posso verificar sua fatura.", + metadata={ + "agent_message_type": "ready", + "expects_user_response": True, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + "wait_timeout_seconds": 60.0, + }, + ) + executor.wait_for_playout_delay_s = 0.03 + runtime.register_callbacks() + + task = asyncio.create_task(runtime.run_pipeline("", "", user_seq=1)) + try: + while not any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands): + await asyncio.sleep(0) + + runtime._session.handlers["user_state_changed"]( + SimpleNamespace(old_state="listening", new_state="speaking") + ) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="ja estou respondendo", is_final=True) + ) + await task + + pipeline_commands = [cmd for cmd in executor.commands if isinstance(cmd, RunPipelineInput)] + self.assertEqual(len(pipeline_commands), 1) + self.assertEqual(runtime.state.user_final_seq, 1) + self.assertTrue( + any( + event == "user_input_dropped" + and fields["agent_message_type"] == "ready" + for event, fields in runtime._timeline.events + ) + ) + self.assertTrue( + any(event == "agent_wait_timeout_armed" for event, _fields in runtime._timeline.events) + ) + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + await self._cancel_pending_tasks() + + async def test_user_state_speaking_during_noninterruptible_presentation_drops_late_final(self) -> None: + runtime, _agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": DEFAULT_WAIT_TEXT, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + _agent.pipeline = _FakeInflightPushPipeline() + runtime.state.current_stage = "PRESENTATION" + runtime.state.current_speech.stage = "PRESENTATION" + runtime.state.current_speech.allow_interruptions = False + runtime.speaking.set() + runtime.register_callbacks() + + runtime._session.handlers["user_state_changed"]( + SimpleNamespace(old_state="listening", new_state="speaking") + ) + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + runtime.speaking.clear() + runtime.state.current_speech.stage = "" + runtime.state.current_speech.allow_interruptions = False + + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="oi", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 0) + self.assertIsNone(runtime.state.pending_user_final) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertFalse(any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands)) + self.assertTrue( + any( + event == "pre_backend_wait_notice_skipped" + and fields["skipped_reason"] == "user_input_drop_pending" + for event, fields in runtime._timeline.events + ) + ) + self.assertTrue( + any( + event == "user_input_dropped" + and fields["agent_message_type"] == "ready" + and fields["text"] == "oi" + for event, fields in runtime._timeline.events + ) + ) + + async def test_user_state_speaking_in_presentation_post_playout_grace_drops_late_final(self) -> None: + runtime, _agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": DEFAULT_WAIT_TEXT, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + _agent.pipeline = _FakeInflightPushPipeline() + runtime.register_callbacks() + + await runtime.say_stage( + "Ola, sou a assistente virtual da TIM.", + "PRESENTATION", + allow_interruptions=False, + schedule_idle_after=False, + ) + + self.assertFalse(runtime.speaking.is_set()) + runtime._session.handlers["user_state_changed"]( + SimpleNamespace(old_state="listening", new_state="speaking") + ) + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="estou falando junto", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 0) + self.assertIsNone(runtime.state.pending_user_final) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertTrue( + any( + event == "user_input_drop_grace_armed" + and fields["grace_ms"] == 250 + for event, fields in runtime._timeline.events + ) + ) + self.assertTrue( + any( + event == "pre_backend_wait_notice_skipped" + and fields["skipped_reason"] == "user_input_drop_pending" + for event, fields in runtime._timeline.events + ) + ) + self.assertTrue( + any( + event == "user_input_dropped" + and fields["reason"] == "user_state_speaking_protected_post_playout_grace" + and fields["agent_message_type"] == "ready" + and fields["text"] == "estou falando junto" + for event, fields in runtime._timeline.events + ) + ) + + async def test_presentation_post_playout_grace_expires_before_next_user_final(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + + with mock.patch( + "app.livekit.runtime.call_runtime.PROTECTED_SPEECH_POST_PLAYOUT_DROP_GRACE_S", + 0.001, + ): + await runtime.say_stage( + "Ola, sou a assistente virtual da TIM.", + "PRESENTATION", + allow_interruptions=False, + schedule_idle_after=False, + ) + await asyncio.sleep(0.01) + + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="agora posso responder", is_final=True) + ) + await self._drain_tasks() + + self.assertEqual(runtime.state.user_final_seq, 1) + self.assertTrue(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_user_state_speaking_interrupts_current_tts(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.current_stage = "ARGUMENTATION" + runtime.state.current_speech.stage = "ARGUMENTATION" + runtime.state.current_speech.allow_interruptions = True + runtime.state.current_speech.handle = executor.speech_handle + runtime.state.current_speech.speech_id = "7f3a0000-0000-4000-8000-000000000000" + runtime.speaking.set() + runtime.register_callbacks() + + runtime._session.handlers["user_state_changed"]( + SimpleNamespace(old_state="listening", new_state="speaking") + ) + await self._drain_tasks() + + interrupt_commands = [ + cmd for cmd in executor.commands if isinstance(cmd, InterruptSpeech) + ] + self.assertEqual(len(interrupt_commands), 1) + self.assertIs(interrupt_commands[0].handle, executor.speech_handle) + + async def test_user_state_speaking_interrupts_inflight_wait_notice(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.current_stage = "AGENT_BACKEND_WAIT" + runtime.state.current_speech.stage = "AGENT_BACKEND_WAIT" + runtime.state.current_speech.allow_interruptions = False + runtime.state.current_speech.handle = executor.speech_handle + runtime.speaking.set() + runtime.register_callbacks() + + runtime._session.handlers["user_state_changed"]( + SimpleNamespace(old_state="listening", new_state="speaking") + ) + await self._drain_tasks() + + interrupt_commands = [ + cmd for cmd in executor.commands if isinstance(cmd, InterruptSpeech) + ] + self.assertEqual(len(interrupt_commands), 1) + self.assertIs(interrupt_commands[0].handle, executor.speech_handle) + self.assertTrue(interrupt_commands[0].force) + + async def test_user_state_speaking_cancels_agent_wait_timeout(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.register_callbacks() + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Ola, como posso ajudar?", + metadata={"wait_timeout_seconds": 0.01}, + ), + user_seq=0, + source="test", + ) + + runtime._session.handlers["user_state_changed"]( + SimpleNamespace(old_state="listening", new_state="speaking") + ) + await asyncio.sleep(0.03) + await self._drain_tasks() + + self.assertFalse(any(isinstance(cmd, NotifyBridgeStop) for cmd in executor.commands)) + self.assertFalse(runtime.finalized.is_set()) + + async def test_user_state_listening_does_not_speak_wait_notice_before_stt_final(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + runtime.register_callbacks() + + runtime._session.handlers["user_state_changed"]( + SimpleNamespace(old_state="speaking", new_state="listening") + ) + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, []) + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_vad_wait_notice_starts_without_waiting_for_user_state_listening(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + runtime.register_callbacks() + + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, [_expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text)]) + + async def test_vad_wait_notice_fast_path_uses_vad_pause_without_user_state(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + "pre_backend_wait_notice_fast_on_vad_pause": True, + } + ) + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime._user_not_speaking.clear() + agent.pipeline = _FakeInflightPushPipeline() + + with mock.patch( + "app.livekit.runtime.call_runtime.asyncio.sleep", + new=mock.AsyncMock(), + ) as sleep: + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await self._drain_tasks() + + sleep.assert_not_awaited() + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, [_expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text)]) + self.assertTrue( + any( + event == "pre_backend_wait_notice_started" + and fields["fast_on_vad_pause"] is True + for event, fields in runtime._timeline.events + ) + ) + + async def test_vad_wait_notice_default_path_skips_while_user_state_speaking(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime._user_not_speaking.clear() + agent.pipeline = _FakeInflightPushPipeline() + + with mock.patch("app.livekit.runtime.call_runtime.PRE_BACKEND_WAIT_NOTICE_GUARD_S", 0): + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await self._drain_tasks() + + self.assertFalse( + any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands) + ) + self.assertTrue( + any( + event == "pre_backend_wait_notice_skipped" + and fields["skipped_reason"] == "user_speaking" + for event, fields in runtime._timeline.events + ) + ) + + def _wait_notice_config(self, wait_text: str) -> dict: + return { + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + + async def test_vad_wait_notice_skipped_while_agent_is_speaking(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides=self._wait_notice_config(wait_text) + ) + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + runtime.state.current_stage = "PRESENTATION" + runtime.state.current_speech.stage = "PRESENTATION" + runtime.speaking.set() + + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await self._drain_tasks() + + self.assertFalse( + any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands) + ) + self.assertTrue( + any( + event == "pre_backend_wait_notice_skipped" + and fields["skipped_reason"] == "agent_speaking" + for event, fields in runtime._timeline.events + ) + ) + self.assertFalse(runtime._pre_backend_wait_notice_reserved) + + async def test_vad_wait_notice_skipped_while_backend_is_processing(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides=self._wait_notice_config(wait_text) + ) + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + await agent._run_lock.acquire() + + try: + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await self._drain_tasks() + finally: + agent._run_lock.release() + + self.assertFalse( + any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands) + ) + self.assertTrue( + any( + event == "pre_backend_wait_notice_skipped" + and fields["skipped_reason"] == "backend_processing" + for event, fields in runtime._timeline.events + ) + ) + self.assertFalse(runtime._pre_backend_wait_notice_reserved) + + async def test_wait_notice_audio_aborts_when_another_speech_took_the_lock(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides=self._wait_notice_config(wait_text) + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + + say_stage_seq_at_schedule = runtime._say_stage_seq + await runtime.say_stage("Resposta do agente.", "PRESENTATION") + + played = await runtime._play_inflight_backend_wait_notice_audio( + user_seq=None, + attempt=1, + abort_if_agent_spoke_since=say_stage_seq_at_schedule, + ) + await self._drain_tasks() + + self.assertFalse(played) + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, ["Resposta do agente."]) + + async def test_vad_wait_notice_does_not_queue_behind_an_agent_speech(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides=self._wait_notice_config(wait_text) + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + + # O aviso e agendado no fim de fala do VAD e so pega o say_lock depois + # que a fala do agente libera: nesse ponto ele nao deve mais sair. + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await runtime.say_stage("Resposta do agente.", "PRESENTATION") + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, ["Resposta do agente."]) + + async def test_vad_wait_notice_returns_reservation_when_it_does_not_play(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides=self._wait_notice_config(wait_text) + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await runtime.say_stage("Resposta do agente.", "PRESENTATION") + await self._drain_tasks() + self.assertFalse(runtime._pre_backend_wait_notice_reserved) + + # A reserva devolvida nao pode bloquear o proximo agendamento. + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual( + speech_texts, + [ + "Resposta do agente.", + _expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text), + ], + ) + + async def test_vad_wait_notice_uses_started_turn_message_id(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, _executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + structured_context = StructuredLogContext( + callid="call-1", + session_id="session-started-turn-runtime", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + runtime._structured_log_context = structured_context + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + reset_turn_message_sequence(structured_context, clear_pending=True) + register_started_turn_message_id(structured_context, "GED-test-0002") + + try: + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await self._drain_tasks() + + wait_events = [ + fields + for event, fields in runtime._timeline.events + if event == "inflight_backend_wait_audio" + ] + self.assertEqual(len(wait_events), 1) + self.assertTrue(wait_events[0]["message_id"].startswith("GED-test-0002_conforto_")) + finally: + reset_turn_message_sequence(structured_context, clear_pending=True) + + async def test_vad_wait_notice_uses_last_agent_message_id_without_started_turn(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, _executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + structured_context = StructuredLogContext( + callid="call-1", + session_id="session-last-agent-turn-runtime", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + runtime._structured_log_context = structured_context + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + reset_turn_message_sequence(structured_context, clear_pending=True) + runtime._remember_agent_base_message_id("MSG-agent-0001") + + try: + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + await self._drain_tasks() + + wait_events = [ + fields + for event, fields in runtime._timeline.events + if event == "inflight_backend_wait_audio" + ] + self.assertEqual(len(wait_events), 1) + self.assertTrue(wait_events[0]["message_id"].startswith("MSG-agent-0001_conforto_")) + finally: + reset_turn_message_sequence(structured_context, clear_pending=True) + + async def test_inflight_backend_wait_notice_skips_post_user_final_grace(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, _executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.last_user_final_at = time.monotonic() + agent.pipeline = _FakeInflightPushPipeline() + + with mock.patch.object( + runtime, + "_wait_post_user_final_grace", + new=mock.AsyncMock(), + ) as wait_grace: + await runtime._play_inflight_backend_wait_notice_audio( + user_seq=None, + attempt=1, + ) + + wait_grace.assert_not_awaited() + + async def test_empty_stt_final_logs_and_clears_pre_backend_wait_notice(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + runtime._pre_backend_wait_notice_reserved = True + + with ( + mock.patch("app.livekit.runtime.call_runtime.log_flow_event") as flow_log, + mock.patch.object(runtime, "arm_idle_timer") as arm_idle, + ): + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="", is_final=True) + ) + await self._drain_tasks() + + self.assertFalse(runtime._pre_backend_wait_notice_reserved) + arm_idle.assert_called_once_with(reason="empty_stt_final") + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertTrue( + any( + call.args[1] == "stt_final_empty" + and call.kwargs["reason"] == "empty_text" + and call.kwargs["pre_backend_wait_reserved"] is True + for call in flow_log.call_args_list + ) + ) + self.assertIn( + ( + "user_transcript_final_empty", + { + "current_stage": "INTRO", + "speaking": False, + "transcript_len": 0, + "pre_backend_wait_reserved": True, + }, + ), + runtime._timeline.events, + ) + + async def test_empty_stt_final_rearms_agent_wait_timeout_when_available(self) -> None: + runtime, _agent, _executor = self._make_runtime() + reply = BackendReply( + stage="PRESENTATION", + text="Ainda esta ai?", + metadata={"wait_timeout_seconds": 60.0}, + ) + runtime.arm_agent_wait_timeout_from_reply(reply, user_seq=0, source="test") + runtime.cancel_agent_wait_timeout("user_state_speaking") + + try: + with ( + mock.patch.object( + runtime, + "_arm_agent_wait_timeout", + wraps=runtime._arm_agent_wait_timeout, + ) as arm_wait, + mock.patch.object(runtime, "arm_idle_timer") as arm_idle, + ): + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="", is_final=True) + ) + + arm_wait.assert_called_once() + self.assertEqual(arm_wait.call_args.kwargs["user_seq"], 0) + self.assertEqual(arm_wait.call_args.kwargs["source"], "empty_stt_final") + arm_idle.assert_not_called() + finally: + await self._cancel_pending_tasks() + + async def test_empty_stt_final_during_comfort_rearms_after_playout(self) -> None: + runtime, _agent, executor = self._make_runtime() + reply = BackendReply( + stage="PRESENTATION", + text="Ainda esta ai?", + metadata={"wait_timeout_seconds": 60.0}, + ) + runtime.arm_agent_wait_timeout_from_reply(reply, user_seq=0, source="test") + runtime.cancel_agent_wait_timeout("user_state_speaking") + executor.wait_for_playout_callbacks.append( + lambda _command: runtime.handle_empty_stt_final(source="stt_provider") + ) + + try: + with mock.patch.object( + runtime, + "_arm_agent_wait_timeout", + wraps=runtime._arm_agent_wait_timeout, + ) as arm_wait: + await runtime.say_stage( + "Um instantinho", + "AGENT_BACKEND_WAIT", + add_to_chat_ctx=False, + allow_interruptions=False, + schedule_idle_after=False, + ) + + arm_wait.assert_called_once() + self.assertEqual(arm_wait.call_args.kwargs["user_seq"], 0) + self.assertEqual(arm_wait.call_args.kwargs["source"], "empty_stt_final") + self.assertFalse(runtime._empty_stt_recovery_pending) + finally: + await self._cancel_pending_tasks() + + async def test_inflight_backend_wait_audio_leaves_tts_metrics_empty_without_exporting_original_id( + self, + ) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime._timeline = _FakeTimeline() + agent.pipeline = _FakeInflightPushPipeline() + executor.speech_handle = SimpleNamespace(id="speech-no-metric", interrupted=False) + structured_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch.dict(os.environ, {"TTS_TTFB_METRIC_WAIT_MS": "0"}), mock.patch( + "app.livekit.runtime.call_runtime.log_structured_event", + side_effect=_capture_event, + ): + await runtime._play_inflight_backend_wait_notice_audio( + user_seq=1, + attempt=2, + message_id="GED-test-0002", + ) + + self.assertEqual(len(structured_events), 1) + self.assertIsNone(structured_events[0]["latencia_total_ms"]) + self.assertIsNone(structured_events[0]["latencia_tffb_ms"]) + self.assertIsNone(structured_events[0]["duracao_audio_ms"]) + self.assertNotIn("original_message_id", structured_events[0]) + + wait_events = [ + fields + for event, fields in runtime._timeline.events + if event == "inflight_backend_wait_audio" + ] + self.assertEqual(len(wait_events), 1) + self.assertTrue(wait_events[0]["message_id"].startswith("GED-test-0002_conforto_")) + self.assertEqual(wait_events[0]["text"], _expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text)) + self.assertNotIn("original_message_id", wait_events[0]) + + async def test_inflight_backend_wait_audio_uses_sidecar_text_for_logged_event(self) -> None: + with TemporaryDirectory() as tmp_dir: + long_dir = Path(tmp_dir) / "long" + long_dir.mkdir() + audio_path = long_dir / "01.wav" + text_path = long_dir / "01.txt" + audio_path.write_bytes(b"fake-wav") + text_path.write_text( + "Estou verificando as informacoes para te ajudar.\nSo um momentinho.", + encoding="utf-8", + ) + expected_text = "Estou verificando as informacoes para te ajudar. So um momentinho." + + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": "Texto padrao que nao deve aparecer.", + "inflight_backend_wait_long_audio_dir": str(long_dir), + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime._timeline = _FakeTimeline() + agent.pipeline = _FakeInflightPushPipeline() + + with ( + mock.patch("app.livekit.runtime.call_runtime.wav_duration_ms", return_value=100), + mock.patch("app.livekit.runtime.call_runtime.wav_audio_frames", return_value="audio-frames"), + ): + await runtime._play_inflight_backend_wait_notice_audio( + user_seq=1, + attempt=2, + message_id="GED-test-0002", + ) + + speech_commands = [ + cmd + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual([cmd.text for cmd in speech_commands], [expected_text]) + wait_events = [ + fields + for event, fields in runtime._timeline.events + if event == "inflight_backend_wait_audio" + ] + self.assertEqual(len(wait_events), 1) + self.assertEqual(wait_events[0]["text"], expected_text) + self.assertEqual(wait_events[0]["path"], str(audio_path)) + + async def test_backend_wait_notice_starts_immediately_after_stt_final(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 0.01, + "inflight_backend_wait_max_notices": 1, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + runtime.state.user_final_seq = 1 + never_finishes = asyncio.Event() + + async def _hang_backend() -> None: + await never_finishes.wait() + + executor.run_pipeline_side_effect = _hang_backend + + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, [_expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text)]) + speech_commands = [ + cmd + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertTrue(speech_commands[0].allow_interruptions) + self.assertIsNotNone(speech_commands[0].audio) + + async def test_pre_stt_wait_notice_activity_controls_next_backend_notice_interval(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 0.01, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + agent.pipeline = _FakeInflightPushPipeline() + runtime._pre_backend_wait_notice_reserved = True + runtime._inflight_backend_activity_at = time.monotonic() - 120.0 + never_finishes = asyncio.Event() + + async def _hang_backend() -> None: + await never_finishes.wait() + + executor.run_pipeline_side_effect = _hang_backend + + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + speech_commands = [ + cmd + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual( + [cmd.text for cmd in speech_commands], + [_expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text)], + ) + self.assertTrue(speech_commands[0].allow_interruptions) + + async def test_inflight_backend_wait_notice_skips_when_user_is_speaking(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + runtime._user_not_speaking.clear() + + await runtime._play_inflight_backend_wait_notice_audio( + user_seq=1, + attempt=1, + message_id="GED-test-0001", + ) + await self._drain_tasks() + + self.assertFalse( + any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands) + ) + + async def test_pre_stt_wait_notice_is_not_restarted_while_pending(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 300.0, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + agent.pipeline = _FakeInflightPushPipeline() + + runtime.schedule_pre_backend_wait_notice(reason="vad_end_of_speech") + runtime.schedule_pre_backend_wait_notice(reason="user_state_listening") + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, [_expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text)]) + + async def test_stt_final_does_not_interrupt_pre_backend_wait_notice(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._pre_backend_wait_notice_reserved = True + runtime.speaking.set() + runtime.state.current_stage = "PRESENTATION" + runtime.state.current_speech.stage = "AGENT_BACKEND_WAIT" + runtime.state.current_speech.allow_interruptions = True + runtime.state.current_speech.handle = executor.speech_handle + + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="Por que veio mais caro?", is_final=True) + ) + await self._drain_tasks() + + self.assertFalse(any(isinstance(cmd, InterruptSpeech) for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_sim_during_pre_backend_wait_notice_runs_pipeline(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._pre_backend_wait_notice_reserved = True + runtime.speaking.set() + runtime.state.current_stage = "PRESENTATION" + runtime.state.current_speech.stage = "AGENT_BACKEND_WAIT" + runtime.state.current_speech.allow_interruptions = False + runtime.state.current_speech.handle = executor.speech_handle + + with mock.patch("app.livekit.runtime.call_runtime.log_flow_event") as flow_log: + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="sim", is_final=True) + ) + await self._drain_tasks() + + self.assertFalse(any(isinstance(cmd, InterruptSpeech) for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertFalse( + any( + call.args[1] == "interrupt_ignored" + and call.kwargs.get("reason") == "backchannel" + and call.kwargs.get("text") == "sim" + for call in flow_log.call_args_list + ) + ) + + async def test_run_pipeline_forwards_pending_speech_interruption(self) -> None: + runtime, agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + agent._consumed_interrupt = ( + "Claro, vou te explicar", + "7f3a0000-0000-4000-8000-000000000000", + False, + ) + + await runtime.run_pipeline("texto", "texto", user_seq=1) + + interruption_commands = [ + cmd for cmd in executor.commands if isinstance(cmd, SetPipelineInterruption) + ] + self.assertEqual(len(interruption_commands), 1) + self.assertTrue(interruption_commands[0].interrupted) + self.assertEqual(interruption_commands[0].listened_text, "Claro, vou te explicar") + self.assertEqual( + interruption_commands[0].speech_id, + "7f3a0000-0000-4000-8000-000000000000", + ) + + async def test_deferred_replay_marks_processing_interruption(self) -> None: + runtime, agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + agent._consumed_interrupt = ( + "", + "7f3a0000-0000-4000-8000-000000000001", + True, + ) + + await runtime.run_pipeline( + "texto interrompido", + "texto interrompido", + user_seq=1, + is_deferred_replay=True, + ) + + command = next( + cmd + for cmd in executor.commands + if isinstance(cmd, SetPipelineProcessingInterruption) + ) + self.assertEqual(command.listened_text, "") + self.assertEqual( + command.speech_id, + "7f3a0000-0000-4000-8000-000000000001", + ) + + async def test_run_pipeline_uses_registered_turn_message_id(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + structured_context = StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + runtime._structured_log_context = structured_context + reset_turn_message_sequence(structured_context, clear_pending=True) + register_transcribed_turn( + structured_context, + message_id="MSG-call-1-0001", + transcription="texto", + text="texto", + ) + + await runtime.run_pipeline("texto", "texto", user_seq=1) + + pipeline_command = next(cmd for cmd in executor.commands if isinstance(cmd, RunPipelineInput)) + self.assertEqual(pipeline_command.user_input["message_id"], "MSG-call-1-0001") + enriched_reply = runtime._backend_reply_with_message_id( + BackendReply(stage="PRESENTATION", text="resposta"), + "MSG-call-1-0001", + ) + self.assertEqual(runtime._message_id_from_reply(enriched_reply), "MSG-call-1-0001") + + async def test_on_user_input_transcribed_clears_started_turn_message_id(self) -> None: + runtime, _agent, executor = self._make_runtime() + structured_context = StructuredLogContext( + callid="call-1", + session_id="session-user-final-started-turn", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + runtime._structured_log_context = structured_context + reset_turn_message_sequence(structured_context, clear_pending=True) + register_started_turn_message_id(structured_context, "MSG-call-1-0002") + register_transcribed_turn( + structured_context, + message_id="MSG-call-1-0002", + transcription="texto", + text="texto", + ) + + try: + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="texto", is_final=True) + ) + await self._drain_tasks() + + pipeline_command = next(cmd for cmd in executor.commands if isinstance(cmd, RunPipelineInput)) + self.assertEqual(pipeline_command.user_input["message_id"], "MSG-call-1-0002") + self.assertEqual(peek_started_turn_message_id(structured_context), "") + finally: + reset_turn_message_sequence(structured_context, clear_pending=True) + + async def test_say_stage_marks_interruption_with_speech_id(self) -> None: + runtime, _agent, executor = self._make_runtime() + executor.speech_handle = SimpleNamespace(interrupted=True) + executor.spoken_text = "Claro, vou te explicar" + + await runtime.say_stage( + "Claro, vou te explicar sua fatura.", + "ARGUMENTATION", + allow_interruptions=True, + speech_id="7f3a0000-0000-4000-8000-000000000000", + ) + + pending_commands = [cmd for cmd in executor.commands if isinstance(cmd, SetPendingInterrupt)] + self.assertEqual(len(pending_commands), 1) + self.assertEqual(pending_commands[0].listened_text, "Claro, vou te explicar") + self.assertEqual( + pending_commands[0].speech_id, + "7f3a0000-0000-4000-8000-000000000000", + ) + + async def test_backend_reply_metadata_controls_tts_interruptibility(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + executor.pipeline_result = BackendReply( + stage="ARGUMENTATION", + text="Ola, como posso ajudar?", + metadata={ + "speech_id": "11111111-1111-4111-8111-111111111111", + "is_interruptible": False, + }, + ) + + await runtime.run_pipeline("texto", "texto", user_seq=1) + + speech_commands = [cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech"] + self.assertEqual(len(speech_commands), 1) + self.assertFalse(speech_commands[0].allow_interruptions) + + async def test_run_pipeline_skips_stale_turn_before_running_pipeline(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 2 + + await runtime.run_pipeline("texto", "texto", user_seq=1) + + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + + async def test_run_pipeline_forwards_short_acknowledgement_while_speaking(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.speaking.set() + + with mock.patch("app.livekit.runtime.call_runtime.log_flow_event") as flow_log: + await runtime.run_pipeline("ok", "ok", user_seq=0) + + self.assertTrue(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertFalse( + any( + call.args[1] == "interrupt_ignored" + and call.kwargs.get("reason") == "backchannel" + for call in flow_log.call_args_list + ) + ) + + async def test_low_confidence_sim_after_tts_runs_pipeline(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + + await runtime.say_stage( + "Consegui esclarecer sua dúvida?", + "PRESENTATION", + allow_interruptions=True, + ) + + self.assertFalse(runtime.speaking.is_set()) + low_confidence_sim = stt_text_with_single_word_threshold( + { + "data": { + "text": "sim", + "words": [{"word": "sim", "probability": 0.009}], + } + }, + min_prob_single_word=0.03, + ) + self.assertEqual(low_confidence_sim, "sim") + + with mock.patch("app.livekit.runtime.call_runtime.log_flow_event") as flow_log: + runtime.on_user_input_transcribed( + SimpleNamespace(transcript=low_confidence_sim, is_final=True) + ) + await self._drain_tasks() + + pipeline_commands = [ + cmd for cmd in executor.commands if isinstance(cmd, RunPipelineInput) + ] + self.assertTrue(pipeline_commands) + self.assertEqual(pipeline_commands[-1].user_input, "sim") + self.assertFalse( + any( + call.args[1] == "interrupt_ignored" + and call.kwargs.get("reason") == "backchannel" + for call in flow_log.call_args_list + ) + ) + + async def test_run_pipeline_drops_output_if_user_speaks_during_pipeline(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + + async def _mutate_seq(): + runtime.state.user_final_seq = 2 + + executor.pipeline_result = BackendReply(stage="PRESENTATION", text="resposta") + executor.run_pipeline_side_effect = _mutate_seq + + await runtime.run_pipeline("texto", "texto", user_seq=1) + + self.assertEqual(runtime.state.current_stage, "INTRO") + self.assertFalse(any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands)) + + async def test_say_stage_waits_briefly_after_user_final(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.last_user_final_at = time.monotonic() + + sleep_calls = [] + + async def _fake_sleep(delay: float) -> None: + sleep_calls.append(delay) + + with mock.patch("app.livekit.runtime.call_runtime.asyncio.sleep", side_effect=_fake_sleep): + await runtime.say_stage("resposta", "ARGUMENTATION") + + self.assertTrue(sleep_calls) + self.assertGreater(sleep_calls[0], 0.0) + self.assertLessEqual(sleep_calls[0], 0.35) + self.assertTrue(any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands)) + + async def test_say_stage_logs_tts_metrics_and_success_status(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._record_xai_tts_turn_timing( + { + "segment_id": "segment-test", + "max_audio_delta_gap_ms": 840, + "max_playout_underrun_0ms": 125, + "xai_micro_underflows": 2, + "xai_avg_underrun_ms": 80, + } + ) + runtime._record_tts_metric( + SimpleNamespace( + metrics=SimpleNamespace( + type="tts_metrics", + ttfb=0.123, + duration=0.456, + audio_duration=1.500, + speech_id="speech-test", + segment_id="segment-test", + label="provider", + ) + ) + ) + structured_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch("app.livekit.runtime.call_runtime.log_structured_event", side_effect=_capture_event): + await runtime.say_stage("resposta", "ARGUMENTATION", message_id="GED-555-0002") + + self.assertTrue(any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands)) + self.assertEqual(len(structured_events), 1) + event = structured_events[0] + self.assertEqual(event["tipo_evento"], "envio msg") + self.assertIsNotNone(event["inicio_ns"]) + self.assertIsNotNone(event["fim_ns"]) + self.assertEqual(event["latencia_total_ms"], 456) + self.assertEqual(event["latencia_tffb_ms"], 123) + self.assertEqual(event["duracao_audio_ms"], 1500) + self.assertEqual(event["tts_max_gap_ms"], 840) + self.assertEqual(event["tts_max_underrun_0ms"], 125) + self.assertEqual(event["tts_underflow_count"], 2) + self.assertEqual(event["tts_avg_underflow_ms"], 80) + self.assertEqual(event["message_id"], "GED-555-0002") + self.assertEqual(event["http_cod_status"], 200) + self.assertEqual(event["http_cod_desc"], "OK") + + async def test_say_stage_leaves_tts_metrics_empty_when_metric_is_missing(self) -> None: + runtime, _agent, executor = self._make_runtime() + executor.speech_handle = SimpleNamespace(id="speech-without-metric", interrupted=False) + structured_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch.dict(os.environ, {"TTS_TTFB_METRIC_WAIT_MS": "0"}), mock.patch( + "app.livekit.runtime.call_runtime.log_structured_event", + side_effect=_capture_event, + ): + await runtime.say_stage("resposta", "ARGUMENTATION", message_id="GED-555-0003") + + self.assertEqual(len(structured_events), 1) + self.assertIsNone(structured_events[0]["latencia_total_ms"]) + self.assertIsNone(structured_events[0]["latencia_tffb_ms"]) + self.assertIsNone(structured_events[0]["duracao_audio_ms"]) + + async def test_say_stage_logs_audio_duration_from_tts_metrics(self) -> None: + runtime, _agent, executor = self._make_runtime() + executor.speech_handle = SimpleNamespace(id="speech-with-audio-metric", interrupted=False) + runtime._record_tts_metric( + SimpleNamespace( + metrics=SimpleNamespace( + type="tts_metrics", + ttfb=0.123, + duration=0.400, + audio_duration=1.500, + speech_id="speech-with-audio-metric", + label="provider", + ) + ) + ) + structured_events = [] + flow_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + def _capture_flow_event(_logger, event_name, **kwargs): + flow_events.append((event_name, kwargs)) + + with mock.patch( + "app.livekit.runtime.call_runtime.log_structured_event", + side_effect=_capture_event, + ), mock.patch( + "app.livekit.runtime.call_runtime.log_flow_event", + side_effect=_capture_flow_event, + ): + await runtime.say_stage("resposta", "ARGUMENTATION", message_id="GED-555-0004") + + self.assertEqual(len(structured_events), 1) + self.assertEqual(structured_events[0]["latencia_total_ms"], 400) + self.assertEqual(structured_events[0]["latencia_tffb_ms"], 123) + self.assertEqual(structured_events[0]["duracao_audio_ms"], 1500) + tts_done_events = [fields for event, fields in flow_events if event == "tts_done"] + self.assertEqual(tts_done_events[0]["duration_ms"], 400) + self.assertEqual(tts_done_events[0]["provider_duration_ms"], 400) + self.assertEqual(tts_done_events[0]["metrics_source"], "livekit_tts_metrics") + + async def test_say_stage_marks_interruption_with_cancelled_tts_metric(self) -> None: + runtime, _agent, executor = self._make_runtime() + executor.speech_handle = SimpleNamespace(id="speech-cancelled", interrupted=True) + executor.spoken_text = "preciso falar" + runtime._record_tts_metric( + SimpleNamespace( + metrics=SimpleNamespace( + type="tts_metrics", + ttfb=0.150, + duration=0.700, + audio_duration=0.250, + speech_id="speech-cancelled", + label="provider", + cancelled=True, + ) + ) + ) + structured_events = [] + flow_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + def _capture_flow_event(_logger, event_name, **kwargs): + flow_events.append((event_name, kwargs)) + + with mock.patch( + "app.livekit.runtime.call_runtime.log_structured_event", + side_effect=_capture_event, + ), mock.patch( + "app.livekit.runtime.call_runtime.log_flow_event", + side_effect=_capture_flow_event, + ): + await runtime.say_stage( + "resposta interrompida", + "ARGUMENTATION", + message_id="GED-555-0005", + allow_interruptions=True, + ) + + self.assertEqual(len(structured_events), 1) + self.assertEqual(structured_events[0]["latencia_total_ms"], 700) + self.assertEqual(structured_events[0]["latencia_tffb_ms"], 150) + self.assertEqual(structured_events[0]["duracao_audio_ms"], 250) + self.assertIs(structured_events[0]["interrupcao"], True) + tts_done_events = [fields for event, fields in flow_events if event == "tts_done"] + self.assertIs(tts_done_events[0]["interruption"], True) + pending_commands = [cmd for cmd in executor.commands if isinstance(cmd, SetPendingInterrupt)] + self.assertEqual(len(pending_commands), 1) + self.assertEqual(pending_commands[0].listened_text, "preciso falar") + + async def test_say_stage_keeps_interruption_duration_without_first_tts_audio(self) -> None: + runtime, _agent, executor = self._make_runtime() + executor.speech_handle = SimpleNamespace(id="speech-cancelled-no-audio", interrupted=True) + runtime._record_tts_metric( + SimpleNamespace( + metrics=SimpleNamespace( + type="tts_metrics", + ttfb=-1.0, + duration=0.700, + audio_duration=0.0, + speech_id="speech-cancelled-no-audio", + label="provider", + cancelled=True, + ) + ) + ) + structured_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch( + "app.livekit.runtime.call_runtime.log_structured_event", + side_effect=_capture_event, + ): + await runtime.say_stage( + "resposta interrompida", + "ARGUMENTATION", + message_id="GED-555-0006", + allow_interruptions=True, + ) + + self.assertEqual(len(structured_events), 1) + self.assertEqual(structured_events[0]["latencia_total_ms"], 700) + self.assertIsNone(structured_events[0]["latencia_tffb_ms"]) + self.assertIsNone(structured_events[0]["duracao_audio_ms"]) + self.assertIs(structured_events[0]["interrupcao"], True) + + async def test_say_stage_returns_to_listening_when_no_first_audio_is_seen(self) -> None: + runtime, _agent, executor = self._make_runtime() + first_handle = SimpleNamespace(id="speech-first", interrupted=False) + executor.speech_handles = [first_handle] + executor.wait_for_playout_delays_s = [0.05] + flow_events = [] + + def capture(_logger, event_name, **kwargs): + flow_events.append((event_name, kwargs)) + + with mock.patch.dict(os.environ, {"TTS_PLAYOUT_START_TIMEOUT_S": "0.01", "TTS_TTFB_METRIC_WAIT_MS": "0"}), mock.patch("app.livekit.runtime.call_runtime.log_flow_event", side_effect=capture): + result = await runtime.say_stage("resposta", "ARGUMENTATION", message_id="GED-555-0005") + + self.assertFalse(result) + speech_commands = [cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech"] + self.assertEqual(len(speech_commands), 1) + self.assertFalse(runtime.finalized.is_set()) + self.assertTrue(first_handle.interrupted) + self.assertTrue(any(event == "tts_failure_return_to_listening" for event, _ in flow_events)) + + async def test_say_stage_returns_to_listening_after_partial_audio_failure(self) -> None: + runtime, _agent, executor = self._make_runtime() + first_handle = SimpleNamespace(id="speech-gap", interrupted=False) + executor.speech_handles = [first_handle] + executor.wait_for_playout_exceptions = [RuntimeError("xAI TTS partial audio failure: underflow_error; socket_resynchronized=0; xai_underrun_estimado_ms=1001; xai_micro_underflows=1; xai_avg_underrun_ms=1001; underflow_error_ms=1000; max_audio_delta_gap_ms=1200; pcm_duration_ms=500; pcm_bytes=24000")] + structured_events = [] + + def capture_event(*_args, **kwargs): + structured_events.append(kwargs) + + with mock.patch.dict(os.environ, {"TTS_TTFB_METRIC_WAIT_MS": "0"}), mock.patch("app.livekit.runtime.call_runtime.log_structured_event", side_effect=capture_event): + result = await runtime.say_stage("resposta com falha no meio", "ARGUMENTATION", message_id="GED-555-0006") + + self.assertFalse(result) + speech_commands = [cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech"] + self.assertEqual(len(speech_commands), 1) + self.assertFalse(runtime.finalized.is_set()) + self.assertTrue(first_handle.interrupted) + self.assertEqual(structured_events[0]["erro_msg"], "Falha TTS") + self.assertIn("underflow_error", structured_events[0]["erro_detalhe"]) + self.assertIn("underflow_error_ms=1000", structured_events[0]["erro_detalhe"]) + self.assertEqual(structured_events[0]["tts_max_gap_ms"], 1200) + self.assertEqual(structured_events[0]["tts_max_underrun_0ms"], 1001) + self.assertEqual(structured_events[0]["tts_underflow_count"], 1) + self.assertEqual(structured_events[0]["tts_avg_underflow_ms"], 1001) + + async def test_say_stage_publishes_partial_failure_without_livekit_metrics(self) -> None: + runtime, _agent, executor = self._make_runtime() + first_handle = SimpleNamespace(id="speech-partial-event", interrupted=False) + executor.speech_handles = [first_handle] + structured_events = [] + + def emit_partial_failure(_command) -> None: + runtime._record_xai_tts_turn_failed( + { + "segment_id": "segment-partial-event", + "reason": "underflow_error", + "max_audio_delta_gap_ms": 920, + "max_playout_underrun_0ms": 1001, + "xai_micro_underflows": 3, + "xai_avg_underrun_ms": 611, + "provider_synthesis_ms": 2450, + "pcm_duration_ms": 860, + } + ) + + executor.wait_for_playout_callbacks = [emit_partial_failure] + + def capture_event(*_args, **kwargs): + structured_events.append(kwargs) + + with mock.patch.dict(os.environ, {"TTS_TTFB_METRIC_WAIT_MS": "0"}), mock.patch( + "app.livekit.runtime.call_runtime.log_structured_event", + side_effect=capture_event, + ): + result = await runtime.say_stage( + "resposta com falha parcial", + "ARGUMENTATION", + message_id="GED-555-0008", + ) + + self.assertFalse(result) + self.assertTrue(first_handle.interrupted) + self.assertEqual(len(structured_events), 1) + event = structured_events[0] + self.assertEqual(event["erro_msg"], "Falha TTS") + self.assertIn("underflow_error", event["erro_detalhe"]) + self.assertEqual(event["latencia_total_ms"], 2450) + self.assertEqual(event["duracao_audio_ms"], 860) + self.assertEqual(event["tts_max_gap_ms"], 920) + self.assertEqual(event["tts_max_underrun_0ms"], 1001) + self.assertEqual(event["tts_underflow_count"], 3) + self.assertEqual(event["tts_avg_underflow_ms"], 611) + self.assertIsNone(event["http_cod_status"]) + + async def test_say_stage_returns_to_listening_after_swallowed_tts_error(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.register_callbacks() + first_handle = SimpleNamespace(id="speech-swallowed", interrupted=False) + executor.speech_handles = [first_handle] + + def emit_error(_command): + error = SimpleNamespace(type="tts_error", error=RuntimeError("no audio frames were pushed for text: resposta"), recoverable=False) + runtime._session.handlers["error"](SimpleNamespace(error=error)) + + executor.wait_for_playout_callbacks = [emit_error] + with mock.patch.dict(os.environ, {"TTS_TTFB_METRIC_WAIT_MS": "0"}): + result = await runtime.say_stage("resposta", "PRESENTATION", message_id="GED-555-0007") + + self.assertFalse(result) + speech_commands = [cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech"] + self.assertEqual(len(speech_commands), 1) + self.assertFalse(runtime.finalized.is_set()) + + async def test_tts_metric_zero_is_not_used_as_ttfb(self) -> None: + runtime, _agent, _executor = self._make_runtime() + + runtime._record_tts_metric( + SimpleNamespace( + metrics=SimpleNamespace( + type="tts_metrics", + ttfb=0, + speech_id="speech-zero", + label="provider", + ) + ) + ) + + self.assertNotIn("speech-zero", runtime._tts_tffb_ms_by_speech_handle_id) + + async def test_finalize_is_idempotent(self) -> None: + runtime, _agent, executor = self._make_runtime() + + await asyncio.gather(runtime.finalize("first"), runtime.finalize("second")) + + self.assertTrue(runtime.finalized.is_set()) + self.assertEqual(sum(isinstance(cmd, EndServiceOnce) for cmd in executor.commands), 1) + self.assertEqual(sum(isinstance(cmd, ExportSession) for cmd in executor.commands), 1) + + async def test_run_pipeline_done_triggers_done_notification_and_finalize(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + executor.pipeline_result = BackendReply(stage="DONE", text="encerrando", done=True) + + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + self.assertTrue(any(isinstance(cmd, NotifyBridgeDone) for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, EndServiceOnce) for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, ExportSession) for cmd in executor.commands)) + + async def test_run_pipeline_non_terminal_reply_does_not_finalize_from_global_done_stage( + self, + ) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + executor.pipeline_result = BackendReply( + stage="PRESENTATION", + text="Perfeito! Seguiremos com o cancelamento.", + done=False, + ) + + def _mark_global_stage_done(_command) -> None: + runtime.state.current_stage = "DONE" + + executor.wait_for_playout_callbacks = [_mark_global_stage_done] + + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + self.assertFalse(any(isinstance(cmd, NotifyBridgeDone) for cmd in executor.commands)) + self.assertFalse(any(isinstance(cmd, EndServiceOnce) for cmd in executor.commands)) + self.assertFalse(runtime.finalized.is_set()) + + async def test_run_pipeline_agent_final_result_speaks_before_terminal_stop(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + executor.pipeline_result = BackendReply( + stage="DONE", + text="Atendimento encerrado com sucesso.", + done=True, + export_payload={ + "type": "nao_resolvido", + "content": "Atendimento encerrado com sucesso.", + "tool_calls": [], + }, + ) + executor.end_reply = BackendReply( + stage="DONE", + text="Nao deveria falar de novo.", + done=True, + export_payload=[{"status": "ok"}], + ) + + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + speech_commands = [cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech"] + stop_commands = [cmd for cmd in executor.commands if isinstance(cmd, NotifyBridgeStop)] + self.assertEqual([cmd.text for cmd in speech_commands], ["Atendimento encerrado com sucesso."]) + self.assertEqual(len(stop_commands), 1) + self.assertEqual(stop_commands[0].status, "stop_nao_resolvido") + self.assertEqual(stop_commands[0].reason, "nao_resolvido") + self.assertEqual(stop_commands[0].phase, "in_session") + self.assertFalse(any(isinstance(cmd, NotifyBridgeDone) for cmd in executor.commands)) + self.assertLess( + executor.commands.index(speech_commands[0]), + executor.commands.index(stop_commands[0]), + ) + self.assertTrue(runtime.finalized.is_set()) + self.assertFalse(any(isinstance(cmd, EndServiceOnce) for cmd in executor.commands)) + self.assertFalse(any(isinstance(cmd, ExportSession) for cmd in executor.commands)) + + async def test_terminal_reply_waits_until_user_stops_speaking(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime._user_not_speaking.clear() + executor.pipeline_result = BackendReply( + stage="DONE", + text="Atendimento encerrado.", + done=True, + export_payload={"type": "nao_resolvido", "content": "Atendimento encerrado."}, + ) + + with mock.patch( + "app.livekit.runtime.call_runtime.POST_USER_FINAL_SPEECH_GRACE_S", + 0.01, + ): + task = asyncio.create_task(runtime.run_pipeline("texto", "texto", user_seq=1)) + await asyncio.sleep(0.01) + self.assertFalse(any(isinstance(cmd, StartSpeech) for cmd in executor.commands)) + + runtime._user_not_speaking.set() + await task + + self.assertEqual( + [cmd.text for cmd in executor.commands if isinstance(cmd, StartSpeech)], + ["Atendimento encerrado."], + ) + + async def test_terminal_reply_is_spoken_when_speaking_becomes_new_user_turn(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime._user_not_speaking.clear() + executor.pipeline_result = BackendReply( + stage="DONE", + text="Atendimento encerrado.", + done=True, + export_payload={"type": "nao_resolvido", "content": "Atendimento encerrado."}, + ) + + with mock.patch( + "app.livekit.runtime.call_runtime.POST_USER_FINAL_SPEECH_GRACE_S", + 0.01, + ): + task = asyncio.create_task(runtime.run_pipeline("texto", "texto", user_seq=1)) + await asyncio.sleep(0.01) + runtime._user_not_speaking.set() + runtime.state.user_final_seq = 2 + await task + + speech_commands = [cmd for cmd in executor.commands if isinstance(cmd, StartSpeech)] + self.assertEqual([cmd.text for cmd in speech_commands], ["Atendimento encerrado."]) + self.assertFalse(speech_commands[0].allow_interruptions) + self.assertTrue(any(isinstance(cmd, NotifyBridgeStop) for cmd in executor.commands)) + self.assertTrue(runtime.finalized.is_set()) + + async def test_run_pipeline_skips_new_turn_after_done_stage(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.current_stage = "DONE" + runtime.state.user_final_seq = 2 + + await runtime.run_pipeline("Ok,", "Ok,", user_seq=2) + + self.assertFalse(any(isinstance(cmd, RunPipelineInput) for cmd in executor.commands)) + self.assertFalse(any(isinstance(cmd, StartSpeech) for cmd in executor.commands)) + + async def test_run_pipeline_ready_wait_timeout_sends_long_silence_stop(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 0 + executor.pipeline_result = BackendReply( + stage="PRESENTATION", + text="Ola, como posso ajudar com sua fatura?", + metadata={"wait_timeout_seconds": 0.01}, + ) + structured_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch("app.livekit.runtime.call_runtime.log_structured_event", side_effect=_capture_event): + await runtime.run_pipeline("", "", user_seq=0) + await self._drain_tasks() + + speech_commands = [cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech"] + stop_commands = [cmd for cmd in executor.commands if isinstance(cmd, NotifyBridgeStop)] + self.assertEqual([cmd.text for cmd in speech_commands], ["Ola, como posso ajudar com sua fatura?"]) + self.assertEqual(len(stop_commands), 1) + self.assertEqual(stop_commands[0].status, "stop_silencio_longo") + self.assertEqual(stop_commands[0].reason, "no_user_response") + self.assertEqual(stop_commands[0].phase, "in_session") + self.assertTrue( + any( + event["tipo_evento"] == "recebimento msg" + and event.get("erro_msg") == "Silencio Longo" + for event in structured_events + ) + ) + self.assertTrue(runtime.finalized.is_set()) + self.assertFalse(any(isinstance(cmd, EndServiceOnce) for cmd in executor.commands)) + self.assertFalse(any(isinstance(cmd, ExportSession) for cmd in executor.commands)) + + async def test_run_pipeline_ready_wait_timeout_sends_retry_messages_before_stop(self) -> None: + controller = _FakeVADThresholdController() + runtime, _agent, executor = self._make_runtime(vad_threshold_controller=controller) + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 0 + retry_messages = [ + "Voce ainda esta na linha?", + "Sigo por aqui.", + "Vou confirmar mais uma vez.", + ] + executor.pipeline_result = BackendReply( + stage="PRESENTATION", + text="Ola, como posso ajudar com sua fatura?", + metadata={ + "message_id": "MSG-timeout-0001", + "wait_timeout_seconds": 0.01, + "wait_retry_messages": retry_messages, + }, + ) + + structured_events = [] + + def _capture_structured_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch( + "app.livekit.runtime.call_runtime.log_structured_event", + side_effect=_capture_structured_event, + ): + await runtime.run_pipeline("", "", user_seq=0) + await self._drain_tasks() + + speech_commands = [cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech"] + stop_commands = [cmd for cmd in executor.commands if isinstance(cmd, NotifyBridgeStop)] + self.assertEqual( + [cmd.text for cmd in speech_commands], + ["Ola, como posso ajudar com sua fatura?", *retry_messages], + ) + self.assertEqual( + [cmd.add_to_chat_ctx for cmd in speech_commands], + [True, False, False, False], + ) + self.assertEqual( + [cmd.allow_interruptions for cmd in speech_commands[1:]], + [True, True, True], + ) + self.assertIn(("activate", 1, "agent_wait_timeout_retry"), controller.calls) + self.assertIn(("restore", "agent_wait_timeout:no_user_response"), controller.calls) + event_ids = [event["message_id"] for event in structured_events] + self.assertEqual( + event_ids[:4], + [ + "MSG-timeout-0001", + "MSG-timeout-0001_inatividade_1", + "MSG-timeout-0001_inatividade_2", + "MSG-timeout-0001_inatividade_3", + ], + ) + self.assertEqual(len(event_ids), 5) + self.assertIn("_erro_terminal_recebimento_", event_ids[4]) + retry_events = [ + fields + for event, fields in runtime._timeline.events + if event == "agent_wait_timeout_retry" + ] + self.assertEqual( + [event["message_id"] for event in retry_events], + [ + "MSG-timeout-0001_inatividade_1", + "MSG-timeout-0001_inatividade_2", + "MSG-timeout-0001_inatividade_3", + ], + ) + self.assertEqual(len(stop_commands), 1) + self.assertEqual(stop_commands[0].status, "stop_silencio_longo") + self.assertEqual(stop_commands[0].reason, "no_user_response") + self.assertEqual(stop_commands[0].phase, "in_session") + self.assertLess( + executor.commands.index(speech_commands[-1]), + executor.commands.index(stop_commands[0]), + ) + self.assertTrue(runtime.finalized.is_set()) + self.assertFalse(any(isinstance(cmd, EndServiceOnce) for cmd in executor.commands)) + self.assertFalse(any(isinstance(cmd, ExportSession) for cmd in executor.commands)) + + async def test_agent_wait_timeout_retry_registers_idle_nudge_event(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 0 + retry_messages = [ + "Voce ainda esta na linha?", + "Sigo por aqui.", + "Vou confirmar mais uma vez.", + ] + executor.pipeline_result = BackendReply( + stage="PRESENTATION", + text="Ola, como posso ajudar com sua fatura?", + metadata={ + "message_id": "MSG-timeout-0001", + "wait_timeout_seconds": 0.01, + "wait_retry_messages": retry_messages, + }, + ) + + await runtime.run_pipeline("", "", user_seq=0) + await self._drain_tasks() + + nudge_commands = [cmd for cmd in executor.commands if isinstance(cmd, InjectIdleNudge)] + self.assertEqual([cmd.text for cmd in nudge_commands], retry_messages) + + # O evento e bufferizado ANTES da fala, para sobreviver a uma interrupcao + # do cliente em cima do retry. + speech_commands = [cmd for cmd in executor.commands if cmd.__class__.__name__ == "StartSpeech"] + for nudge, speech in zip(nudge_commands, speech_commands[1:]): + self.assertLess( + executor.commands.index(nudge), + executor.commands.index(speech), + ) + + async def test_agent_wait_timeout_metadata_persists_for_later_replies(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + + try: + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Ola, como posso ajudar?", + metadata={ + "wait_timeout_seconds": 60.0, + "wait_retry_messages": ["Voce ainda esta na linha?"], + }, + ), + user_seq=0, + source="initial", + ) + runtime.cancel_agent_wait_timeout("user_final") + runtime._clear_agent_wait_timeout_memory() + runtime.state.user_final_seq = 1 + + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Pode me responder?", + ), + user_seq=1, + source="next_reply", + ) + + armed_events = [ + fields + for event, fields in runtime._timeline.events + if event == "agent_wait_timeout_armed" + ] + self.assertEqual(len(armed_events), 2) + self.assertEqual(armed_events[-1]["wait_timeout_s"], 60.0) + self.assertEqual(armed_events[-1]["retry_messages"], 1) + self.assertEqual(armed_events[-1]["source"], "next_reply") + self.assertEqual(armed_events[-1]["user_seq"], 1) + finally: + await self._cancel_pending_tasks() + + async def test_agent_wait_timeout_metadata_rewrites_persisted_config(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + + try: + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Primeira pergunta", + metadata={ + "wait_timeout_seconds": 60.0, + "wait_retry_messages": ["Mensagem antiga"], + }, + ), + user_seq=0, + source="initial", + ) + runtime.cancel_agent_wait_timeout("user_final") + runtime._clear_agent_wait_timeout_memory() + runtime.state.user_final_seq = 1 + + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Nova pergunta", + metadata={ + "wait_timeout_seconds": 30.0, + "wait_retry_messages": ["Nova mensagem", "Ultimo aviso"], + }, + ), + user_seq=1, + source="rewrite", + ) + runtime.cancel_agent_wait_timeout("user_final") + runtime._clear_agent_wait_timeout_memory() + runtime.state.user_final_seq = 2 + + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Pergunta sem metadata", + ), + user_seq=2, + source="after_rewrite", + ) + + armed_events = [ + fields + for event, fields in runtime._timeline.events + if event == "agent_wait_timeout_armed" + ] + self.assertEqual(len(armed_events), 3) + self.assertEqual(armed_events[-2]["wait_timeout_s"], 30.0) + self.assertEqual(armed_events[-2]["retry_messages"], 2) + self.assertEqual(armed_events[-1]["wait_timeout_s"], 30.0) + self.assertEqual(armed_events[-1]["retry_messages"], 2) + self.assertEqual(armed_events[-1]["source"], "after_rewrite") + self.assertEqual(armed_events[-1]["user_seq"], 2) + finally: + await self._cancel_pending_tasks() + + async def test_feedback_reply_does_not_arm_client_wait_timeout(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + + try: + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Posso verificar sua fatura.", + metadata={ + "agent_message_type": "ready", + "expects_user_response": True, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + "wait_timeout_seconds": 60.0, + "wait_retry_messages": ["Voce esta ai?"], + }, + ), + user_seq=0, + source="ready", + ) + runtime.cancel_agent_wait_timeout("test") + runtime._clear_agent_wait_timeout_memory() + + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Ainda estou consultando sua fatura.", + metadata={ + "agent_message_type": "feedback", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + "event": "feedback", + }, + ), + user_seq=0, + source="feedback", + ) + + armed_events = [ + fields + for event, fields in runtime._timeline.events + if event == "agent_wait_timeout_armed" + ] + self.assertEqual(len(armed_events), 1) + + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Consegui consultar. Posso te ajudar em mais algo?", + metadata={ + "agent_message_type": "result", + "agent_result_type": "final", + "expects_user_response": True, + "drop_user_input_while_speaking": False, + }, + ), + user_seq=0, + source="final", + ) + + armed_events = [ + fields + for event, fields in runtime._timeline.events + if event == "agent_wait_timeout_armed" + ] + self.assertEqual(len(armed_events), 2) + self.assertEqual(armed_events[-1]["source"], "final") + self.assertEqual(armed_events[-1]["retry_messages"], 1) + finally: + await self._cancel_pending_tasks() + + async def test_agent_wait_timeout_zero_metadata_clears_persisted_config(self) -> None: + runtime, _agent, _executor = self._make_runtime() + runtime._timeline = _FakeTimeline() + + try: + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Primeira pergunta", + metadata={"wait_timeout_seconds": 60.0}, + ), + user_seq=0, + source="initial", + ) + + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Sem timeout daqui em diante", + metadata={"wait_timeout_seconds": 0}, + ), + user_seq=0, + source="disable", + ) + runtime.state.user_final_seq = 1 + runtime.arm_agent_wait_timeout_from_reply( + BackendReply( + stage="PRESENTATION", + text="Pergunta sem metadata", + ), + user_seq=1, + source="after_disable", + ) + + armed_events = [ + fields + for event, fields in runtime._timeline.events + if event == "agent_wait_timeout_armed" + ] + self.assertEqual(len(armed_events), 1) + finally: + await self._cancel_pending_tasks() + + async def test_run_pipeline_backend_failure_triggers_terminal_stop_and_finalize(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + + async def _fail_backend() -> None: + raise RuntimeError("backend closed connection") + + executor.run_pipeline_side_effect = _fail_backend + structured_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch("app.livekit.runtime.call_runtime.log_structured_event", side_effect=_capture_event): + await runtime.run_pipeline("texto", "texto", user_seq=1) + + stop_commands = [cmd for cmd in executor.commands if isinstance(cmd, NotifyBridgeStop)] + self.assertEqual(len(stop_commands), 1) + self.assertEqual(stop_commands[0].status, "stop_agent_backend_unavailable") + self.assertEqual(stop_commands[0].reason, "resource_unhealthy") + self.assertEqual(stop_commands[0].resource, "agent_backend") + self.assertEqual(stop_commands[0].failed_resources, ("agent_backend",)) + self.assertTrue( + any( + event["tipo_evento"] == "envio msg" + and event.get("erro_msg") == "Falha comunicacao" + and "_erro_terminal_envio_" in event.get("message_id", "") + for event in structured_events + ) + ) + self.assertTrue(any(isinstance(cmd, EndServiceOnce) for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, ExportSession) for cmd in executor.commands)) + + async def test_run_pipeline_transfer_reply_uses_suffixed_message_id(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime.state.user_final_seq = 1 + executor.pipeline_result = BackendReply( + stage="PRESENTATION", + text="", + done=False, + export_payload={ + "type": "transferred", + "status": "transferred", + "additionalInformations": {"changeAgentId": "ATH-001"}, + }, + ) + structured_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch("app.livekit.runtime.call_runtime.log_structured_event", side_effect=_capture_event): + await runtime.run_pipeline( + "texto", + "texto", + user_seq=1, + message_id="GED-test-0002", + ) + + self.assertTrue( + any( + event["tipo_evento"] == "envio msg" + and event.get("erro_msg") == "Transferido" + and event.get("message_id", "").startswith("GED-test-0002_transferencia_") + for event in structured_events + ) + ) + + async def test_finalize_speaks_end_reply_when_remote_listener_exists(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + executor.end_reply = BackendReply( + stage="DONE", + text="Mensagem final do backend.", + done=True, + export_payload=[{"status": "ok"}], + ) + + await runtime.finalize("manual") + + self.assertTrue(any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, ExportSession) for cmd in executor.commands)) + + async def test_finalize_skips_end_reply_tts_without_remote_listener(self) -> None: + runtime, _agent, executor = self._make_runtime() + executor.end_reply = BackendReply( + stage="DONE", + text="Mensagem final do backend.", + done=True, + export_payload=[{"status": "ok"}], + ) + + await runtime.finalize("manual") + + self.assertFalse(any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands)) + self.assertTrue(any(isinstance(cmd, ExportSession) for cmd in executor.commands)) + + async def test_register_callbacks_wires_session_and_shutdown(self) -> None: + runtime, _agent, _executor = self._make_runtime() + + runtime.register_callbacks() + + self.assertIn("user_input_transcribed", runtime._session.handlers) + self.assertIn("close", runtime._session.handlers) + self.assertIn("data_received", runtime._ctx.room.handlers) + self.assertEqual(len(runtime._ctx.shutdown_callbacks), 1) + + async def test_session_close_error_triggers_resource_stop(self) -> None: + runtime, _agent, executor = self._make_runtime() + + runtime.register_callbacks() + runtime._session.handlers["close"]( + SimpleNamespace( + reason=SimpleNamespace(value="error"), + error=SimpleNamespace(type="stt_error"), + ) + ) + await self._drain_tasks() + + stop_commands = [cmd for cmd in executor.commands if isinstance(cmd, NotifyBridgeStop)] + self.assertEqual(len(stop_commands), 1) + self.assertEqual(stop_commands[0].status, "stop_stt_unavailable") + self.assertEqual(stop_commands[0].resource, "stt") + + async def test_bridge_control_client_audio_enabled_triggers_initial_agent_turn(self) -> None: + runtime, _agent, executor = self._make_runtime(agent_starts_conversation=True) + runtime._ctx.room.remote_participants = {"user-1": object()} + executor.pipeline_result = BackendReply(stage="PRESENTATION", text="Ola, sou o agent.") + + runtime.register_callbacks() + runtime._ctx.room.handlers["data_received"]( + SimpleNamespace( + topic="bridge.control", + data=b'{"type":"client_audio_enabled","room":"room-test","protocol":"PRT-1"}', + ) + ) + await self._drain_tasks() + + run_commands = [cmd for cmd in executor.commands if isinstance(cmd, RunPipelineInput)] + self.assertEqual(len(run_commands), 1) + self.assertEqual(run_commands[0].user_input, "") + self.assertTrue(any(cmd.__class__.__name__ == "StartSpeech" for cmd in executor.commands)) + + async def test_run_holds_user_audio_input_before_starting_session(self) -> None: + runtime, _agent, executor = self._make_runtime(agent_starts_conversation=True) + runtime._ctx.room.remote_participants = {"user-1": object()} + audio_enabled_on_start = [] + executor.on_start_session = lambda _cmd: audio_enabled_on_start.append( + runtime._session.input.audio_enabled + ) + + await runtime.run() + await asyncio.sleep(0) + await self._cancel_pending_tasks() + + # O RoomIO le o estado do input ao anexar a track: se o hold viesse depois + # do StartSession, o audio do setup ja teria chegado no VAD/STT. + self.assertEqual(audio_enabled_on_start, [False]) + self.assertFalse(runtime._session.input.audio_enabled) + self.assertEqual(runtime._session.input.calls, [False]) + + async def test_run_keeps_user_audio_input_when_user_starts_conversation(self) -> None: + runtime, _agent, _executor = self._make_runtime(agent_starts_conversation=False) + runtime._ctx.room.remote_participants = {"user-1": object()} + + await runtime.run() + await asyncio.sleep(0) + await self._cancel_pending_tasks() + + self.assertTrue(runtime._session.input.audio_enabled) + self.assertEqual(runtime._session.input.calls, []) + + async def test_initial_agent_turn_releases_user_audio_input_after_playout(self) -> None: + runtime, _agent, executor = self._make_runtime(agent_starts_conversation=True) + runtime._ctx.room.remote_participants = {"user-1": object()} + executor.pipeline_result = BackendReply(stage="PRESENTATION", text="Ola, sou o agent.") + audio_enabled_during_playout = [] + executor.wait_for_playout_callbacks = [ + lambda _cmd: audio_enabled_during_playout.append( + runtime._session.input.audio_enabled + ) + ] + + runtime.hold_user_audio_input(reason="room_setup") + runtime.register_callbacks() + runtime._ctx.room.handlers["data_received"]( + SimpleNamespace( + topic="bridge.control", + data=b'{"type":"client_audio_enabled","room":"room-test","protocol":"PRT-1"}', + ) + ) + await self._drain_tasks() + + self.assertEqual(audio_enabled_during_playout, [False]) + self.assertTrue(runtime._session.input.audio_enabled) + self.assertEqual(runtime._session.input.calls, [False, True]) + + async def test_initial_agent_turn_releases_user_audio_input_when_pipeline_fails(self) -> None: + runtime, _agent, executor = self._make_runtime(agent_starts_conversation=True) + runtime._ctx.room.remote_participants = {"user-1": object()} + + async def _boom() -> None: + raise RuntimeError("backend fora do ar") + + executor.run_pipeline_side_effect = _boom + + runtime.hold_user_audio_input(reason="room_setup") + runtime.register_callbacks() + runtime._ctx.room.handlers["data_received"]( + SimpleNamespace( + topic="bridge.control", + data=b'{"type":"client_audio_enabled","room":"room-test","protocol":"PRT-1"}', + ) + ) + await self._drain_tasks() + + # Falha no turno inicial nao pode deixar a chamada surda. + self.assertTrue(runtime._session.input.audio_enabled) + + async def test_user_audio_input_gate_timeout_releases_when_initial_turn_never_starts( + self, + ) -> None: + runtime, _agent, _executor = self._make_runtime(agent_starts_conversation=True) + runtime._ctx.room.remote_participants = {"user-1": object()} + + with mock.patch.dict(os.environ, {"USER_AUDIO_INPUT_SETUP_TIMEOUT_S": "0.01"}): + runtime.hold_user_audio_input(reason="room_setup") + runtime._arm_user_audio_input_gate_timeout() + await self._drain_tasks() + + # client_audio_enabled nunca chegou: o gate abre sozinho. + self.assertTrue(runtime._session.input.audio_enabled) + self.assertEqual(runtime._session.input.calls, [False, True]) + + async def test_backend_push_message_is_spoken_without_new_user_turn(self) -> None: + runtime, _agent, executor = self._make_runtime( + agent_starts_conversation=True, + push_replies=[BackendReply(stage="PRESENTATION", text="Segunda mensagem do backend.")], + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + executor.pipeline_result = BackendReply(stage="PRESENTATION", text="Primeira mensagem do backend.") + + runtime.register_callbacks() + runtime._ctx.room.handlers["data_received"]( + SimpleNamespace( + topic="bridge.control", + data=b'{"type":"client_audio_enabled","room":"room-test","protocol":"PRT-1"}', + ) + ) + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual( + speech_texts, + [ + "Primeira mensagem do backend.", + "Segunda mensagem do backend.", + ], + ) + + async def test_inflight_backend_feedback_is_spoken_before_final_reply(self) -> None: + runtime, agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + + pipeline = _FakeInflightPushPipeline() + agent.pipeline = pipeline + executor.pipeline_result = BackendReply( + stage="PRESENTATION", + text="Resposta final do backend.", + ) + + async def _emit_feedback_then_finish() -> None: + await pipeline.push( + BackendReply( + stage="PRESENTATION", + text="Ainda estou consultando sua fatura.", + metadata={ + "event": "feedback", + "agent_message_type": "feedback", + "expects_user_response": False, + "drop_user_input_while_speaking": True, + "is_interruptible": False, + }, + ) + ) + await asyncio.sleep(0.01) + pipeline.close() + + executor.run_pipeline_side_effect = _emit_feedback_then_finish + + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual( + speech_texts, + [ + "Ainda estou consultando sua fatura.", + "Resposta final do backend.", + ], + ) + speech_commands = [ + cmd + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertFalse(speech_commands[0].allow_interruptions) + + async def test_inflight_backend_wait_notice_is_spoken_immediately_after_stt_final(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 0.01, + "inflight_backend_wait_max_notices": 1, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + agent.pipeline = _FakeInflightPushPipeline() + never_finishes = asyncio.Event() + + async def _hang_backend() -> None: + await never_finishes.wait() + + executor.run_pipeline_side_effect = _hang_backend + + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, [_expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text)]) + + async def test_inflight_backend_wait_audio_sequence_starts_short_then_long(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + with TemporaryDirectory() as tmp_dir: + base_dir = Path(tmp_dir) + short_dir = base_dir / "short" + long_dir = base_dir / "long" + short_dir.mkdir() + long_dir.mkdir() + short_audio = short_dir / "curto.wav" + long_audio = long_dir / "longo.wav" + short_audio.write_bytes(b"short") + long_audio.write_bytes(b"long") + + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 0.005, + "inflight_backend_wait_timeout_s": 0.6, + "inflight_backend_wait_max_notices": 2, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_short_audio_dir": str(short_dir), + "inflight_backend_wait_long_audio_dir": str(long_dir), + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime._timeline = _FakeTimeline() + agent.pipeline = _FakeInflightPushPipeline() + never_finishes = asyncio.Event() + + async def _hang_backend() -> None: + await never_finishes.wait() + + def _audio_frames(path: str): + return f"audio:{Path(path).name}" + + executor.run_pipeline_side_effect = _hang_backend + + with ( + mock.patch("app.livekit.runtime.call_runtime.wav_duration_ms", return_value=100), + mock.patch( + "app.livekit.runtime.call_runtime.wav_audio_frames", + side_effect=_audio_frames, + ), + mock.patch.object(runtime, "_log_resource_error_events") as structured_error, + ): + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + speech_commands = [ + cmd + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual( + [cmd.audio for cmd in speech_commands], + ["audio:curto.wav", "audio:longo.wav"], + ) + max_notice_events = [ + fields + for event, fields in runtime._timeline.events + if event == "inflight_backend_wait_max_notices_reached" + ] + self.assertEqual(len(max_notice_events), 1) + self.assertEqual(max_notice_events[0]["notices_sent"], 2) + self.assertEqual(max_notice_events[0]["max_notices"], 2) + max_notice_structured_calls = [ + call + for call in structured_error.call_args_list + if call.kwargs.get("reason") + == "inflight_backend_wait_max_notices_reached" + ] + self.assertEqual(len(max_notice_structured_calls), 1) + self.assertEqual( + max_notice_structured_calls[0].kwargs["resource"], + "agent_backend", + ) + + async def test_inflight_backend_wait_zero_max_notices_is_unlimited(self) -> None: + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 0.005, + "inflight_backend_wait_timeout_s": 0.2, + "inflight_backend_wait_max_notices": 0, + "inflight_backend_wait_text": DEFAULT_WAIT_TEXT, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime._timeline = _FakeTimeline() + agent.pipeline = _FakeInflightPushPipeline() + backend_done = asyncio.Event() + attempts: list[int] = [] + + async def _hang_backend() -> None: + await backend_done.wait() + + async def _play_notice(**kwargs): + attempts.append(kwargs["attempt"]) + if len(attempts) == 3: + backend_done.set() + return InflightBackendWaitNoticeResult(started=True) + + executor.run_pipeline_side_effect = _hang_backend + + with mock.patch.object( + runtime, + "_play_inflight_backend_wait_notice_audio", + side_effect=_play_notice, + ): + await runtime.run_pipeline("texto", "texto", user_seq=1) + + self.assertEqual(attempts, [1, 2, 3]) + self.assertFalse( + any( + event == "inflight_backend_wait_max_notices_reached" + for event, _fields in runtime._timeline.events + ) + ) + + async def test_inflight_backend_wait_skipped_notice_does_not_consume_attempt(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + with TemporaryDirectory() as tmp_dir: + base_dir = Path(tmp_dir) + short_dir = base_dir / "short" + short_dir.mkdir() + short_audio = short_dir / "curto.wav" + short_audio.write_bytes(b"short") + + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 0.005, + "inflight_backend_wait_timeout_s": 0.06, + "inflight_backend_wait_max_notices": 1, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_short_audio_dir": str(short_dir), + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime._user_not_speaking.clear() + agent.pipeline = _FakeInflightPushPipeline() + never_finishes = asyncio.Event() + + async def _hang_backend() -> None: + await never_finishes.wait() + + async def _mark_user_not_speaking() -> None: + await asyncio.sleep(0.015) + runtime._user_not_speaking.set() + + def _audio_frames(path: str): + return f"audio:{Path(path).name}" + + executor.run_pipeline_side_effect = _hang_backend + restore_task = asyncio.create_task(_mark_user_not_speaking()) + + try: + with ( + mock.patch("app.livekit.runtime.call_runtime.wav_duration_ms", return_value=100), + mock.patch( + "app.livekit.runtime.call_runtime.wav_audio_frames", + side_effect=_audio_frames, + ), + ): + await runtime.run_pipeline("texto", "texto", user_seq=1) + finally: + restore_task.cancel() + await asyncio.gather(restore_task, return_exceptions=True) + + await self._drain_tasks() + + speech_commands = [ + cmd + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual([cmd.audio for cmd in speech_commands], ["audio:curto.wav"]) + + async def test_inflight_backend_wait_interval_counts_after_audio_playout(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 0.025, + "inflight_backend_wait_timeout_s": 0.04, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + agent.pipeline = _FakeInflightPushPipeline() + executor.wait_for_playout_delay_s = 0.02 + never_finishes = asyncio.Event() + + async def _hang_backend() -> None: + await never_finishes.wait() + + executor.run_pipeline_side_effect = _hang_backend + + await runtime.run_pipeline("texto", "texto", user_seq=1) + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + self.assertEqual(speech_texts, [_expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text)]) + + async def test_short_processing_interruption_rearms_periodic_notice_deadline(self) -> None: + interval_s = 0.05 + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": interval_s, + "inflight_backend_wait_timeout_s": 0.3, + "inflight_backend_wait_max_notices": 2, + "inflight_backend_wait_text": DEFAULT_WAIT_TEXT, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + agent.pipeline = _FakeInflightPushPipeline() + backend_done = asyncio.Event() + attempts: list[int] = [] + attempted_at: list[float] = [] + + async def _hang_backend() -> None: + await backend_done.wait() + + async def _finish_short_speech() -> None: + await asyncio.sleep(0.01) + runtime.note_vad_speech_end(300) + runtime._user_not_speaking.set() + + async def _play_notice(**kwargs): + attempt = kwargs["attempt"] + attempts.append(attempt) + attempted_at.append(time.monotonic()) + if attempts == [1]: + return InflightBackendWaitNoticeResult(started=True) + if attempts == [1, 2]: + runtime._user_not_speaking.clear() + asyncio.create_task(_finish_short_speech()) + return InflightBackendWaitNoticeResult( + started=True, + interrupted_by_user=True, + ) + backend_done.set() + return InflightBackendWaitNoticeResult(started=True) + + executor.run_pipeline_side_effect = _hang_backend + + with mock.patch.object( + runtime, + "_play_inflight_backend_wait_notice_audio", + side_effect=_play_notice, + ): + await runtime.run_pipeline("texto", "texto", user_seq=1) + + self.assertEqual(attempts, [1, 2, 2]) + self.assertGreaterEqual(attempted_at[2] - attempted_at[1], interval_s) + + async def test_short_speech_during_silent_backend_wait_does_not_rearm_notice_deadline(self) -> None: + interval_s = 0.05 + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": interval_s, + "inflight_backend_wait_timeout_s": 0.3, + "inflight_backend_wait_max_notices": 2, + "inflight_backend_wait_text": DEFAULT_WAIT_TEXT, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + agent.pipeline = _FakeInflightPushPipeline() + backend_done = asyncio.Event() + attempts: list[int] = [] + attempted_at: list[float] = [] + short_speech_ended_at = 0.0 + + async def _hang_backend() -> None: + await backend_done.wait() + + async def _emit_short_speech_during_wait() -> None: + nonlocal short_speech_ended_at + await asyncio.sleep(interval_s / 2) + runtime.note_vad_speech_end(300) + short_speech_ended_at = time.monotonic() + + async def _play_notice(**kwargs): + attempts.append(kwargs["attempt"]) + attempted_at.append(time.monotonic()) + if attempts == [1]: + asyncio.create_task(_emit_short_speech_during_wait()) + else: + backend_done.set() + return InflightBackendWaitNoticeResult(started=True) + + executor.run_pipeline_side_effect = _hang_backend + + with mock.patch.object( + runtime, + "_play_inflight_backend_wait_notice_audio", + side_effect=_play_notice, + ): + await runtime.run_pipeline("texto", "texto", user_seq=1) + + self.assertEqual(attempts, [1, 2]) + self.assertGreater(short_speech_ended_at, 0.0) + self.assertLess(attempted_at[1] - short_speech_ended_at, interval_s) + + async def test_valid_processing_interruption_rearms_periodic_notice_deadline(self) -> None: + interval_s = 0.04 + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": interval_s, + "inflight_backend_wait_timeout_s": 0.3, + "inflight_backend_wait_max_notices": 2, + "inflight_backend_wait_text": DEFAULT_WAIT_TEXT, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + agent.pipeline = _FakeInflightPushPipeline() + backend_done = asyncio.Event() + attempts: list[int] = [] + attempted_at: list[float] = [] + + async def _hang_backend() -> None: + await backend_done.wait() + + async def _finish_valid_speech() -> None: + await asyncio.sleep(0.01) + runtime.note_vad_speech_end(1000) + runtime._user_not_speaking.set() + + async def _play_notice(**kwargs): + attempt = kwargs["attempt"] + attempts.append(attempt) + attempted_at.append(time.monotonic()) + if attempts == [1]: + return InflightBackendWaitNoticeResult(started=True) + if attempts == [1, 2]: + runtime._user_not_speaking.clear() + asyncio.create_task(_finish_valid_speech()) + return InflightBackendWaitNoticeResult( + started=True, + interrupted_by_user=True, + ) + backend_done.set() + return InflightBackendWaitNoticeResult(started=True) + + executor.run_pipeline_side_effect = _hang_backend + + with ( + mock.patch.object( + runtime, + "_play_inflight_backend_wait_notice_audio", + side_effect=_play_notice, + ), + mock.patch.object(runtime, "_play_deferred_interruption_comfort"), + ): + await runtime.run_pipeline("texto", "texto", user_seq=1) + + self.assertEqual(attempts, [1, 2, 2]) + self.assertGreaterEqual( + attempted_at[2] - attempted_at[1], + interval_s, + ) + + async def test_long_processing_interruption_keeps_periodic_notices_active(self) -> None: + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 0.01, + "inflight_backend_wait_timeout_s": 0.2, + "inflight_backend_wait_max_notices": 2, + "inflight_backend_wait_text": DEFAULT_WAIT_TEXT, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime.state.deferred_interruption.backend_in_flight = True + agent.pipeline = _FakeInflightPushPipeline() + backend_done = asyncio.Event() + attempts: list[int] = [] + + async def _hang_backend() -> None: + await backend_done.wait() + + async def _interrupt_after_first_notice() -> None: + await asyncio.sleep(0) + runtime.note_vad_speech_end(1500) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="e nao resolve", is_final=True) + ) + + async def _play_notice(**kwargs): + attempts.append(kwargs["attempt"]) + if attempts == [1]: + asyncio.create_task(_interrupt_after_first_notice()) + else: + backend_done.set() + return InflightBackendWaitNoticeResult(started=True) + + executor.run_pipeline_side_effect = _hang_backend + + with ( + mock.patch.object( + runtime, + "_play_inflight_backend_wait_notice_audio", + side_effect=_play_notice, + ), + mock.patch.object(runtime, "_play_deferred_interruption_comfort"), + ): + await runtime._execute_pipeline_with_inflight_backend_wait( + "pedido original", + user_seq=1, + ) + + self.assertEqual(runtime.state.user_final_seq, 2) + self.assertEqual( + [turn.text for turn in runtime.state.deferred_interruption.long_turns], + ["e nao resolve"], + ) + self.assertEqual(attempts, [1, 2]) + + async def test_long_processing_interruption_does_not_disable_backend_timeout(self) -> None: + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 60.0, + "inflight_backend_wait_timeout_s": 0.03, + "inflight_backend_wait_max_notices": 1, + "inflight_backend_wait_text": DEFAULT_WAIT_TEXT, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + runtime.state.deferred_interruption.backend_in_flight = True + agent.pipeline = _FakeInflightPushPipeline() + never_finishes = asyncio.Event() + + async def _hang_backend() -> None: + await never_finishes.wait() + + async def _play_notice(**_kwargs): + runtime.note_vad_speech_end(1500) + runtime.on_user_input_transcribed( + SimpleNamespace(transcript="continua demorando", is_final=True) + ) + return InflightBackendWaitNoticeResult(started=True) + + executor.run_pipeline_side_effect = _hang_backend + + with ( + mock.patch.object( + runtime, + "_play_inflight_backend_wait_notice_audio", + side_effect=_play_notice, + ), + mock.patch.object(runtime, "_play_deferred_interruption_comfort"), + ): + with self.assertRaises(InflightBackendWaitTimedOut): + await runtime._execute_pipeline_with_inflight_backend_wait( + "pedido original", + user_seq=1, + ) + + self.assertEqual(runtime.state.user_final_seq, 2) + self.assertTrue(runtime.finalized.is_set()) + + async def test_inflight_backend_wait_notice_then_backend_unavailable_stop(self) -> None: + wait_text = "Um momento, ainda estou consultando para te ajudar." + runtime, agent, executor = self._make_runtime( + runtime_config_overrides={ + "inflight_backend_wait_interval_s": 0.01, + "inflight_backend_wait_timeout_s": 0.035, + "inflight_backend_wait_max_notices": 6, + "inflight_backend_wait_text": wait_text, + "inflight_backend_wait_long_audio_path": WAIT_LONG_AUDIO_PATH, + } + ) + runtime._ctx.room.remote_participants = {"user-1": object()} + runtime.state.user_final_seq = 1 + agent.pipeline = _FakeInflightPushPipeline() + never_finishes = asyncio.Event() + + async def _hang_backend() -> None: + await never_finishes.wait() + + executor.run_pipeline_side_effect = _hang_backend + + structured_events = [] + + def _capture_event(*args, **kwargs): + structured_events.append(kwargs) + + with mock.patch("app.livekit.runtime.call_runtime.log_structured_event", side_effect=_capture_event): + await runtime.run_pipeline( + "texto", + "texto", + user_seq=1, + message_id="GED-test-0002", + ) + await self._drain_tasks() + + speech_texts = [ + cmd.text + for cmd in executor.commands + if cmd.__class__.__name__ == "StartSpeech" + ] + stop_commands = [cmd for cmd in executor.commands if isinstance(cmd, NotifyBridgeStop)] + + self.assertGreaterEqual(len(speech_texts), 1) + self.assertTrue( + all( + text == _expected_wait_text_for_audio(WAIT_LONG_AUDIO_PATH, wait_text) + for text in speech_texts + ) + ) + comfort_events = [ + event + for event in structured_events + if str(event.get("message_id") or "").startswith("GED-test-0002_conforto_") + ] + self.assertGreaterEqual(len(comfort_events), 1) + self.assertTrue( + any( + event["tipo_evento"] == "envio msg" + and event.get("erro_msg") == "Falha comunicacao" + and "_erro_terminal_envio_" in event.get("message_id", "") + for event in structured_events + ) + ) + self.assertEqual(len(stop_commands), 1) + self.assertEqual(stop_commands[0].status, "stop_agent_backend_unavailable") + self.assertEqual(stop_commands[0].reason, "resource_unhealthy") + self.assertEqual(stop_commands[0].resource, "agent_backend") + self.assertEqual(stop_commands[0].failed_resources, ("agent_backend",)) + self.assertTrue(runtime.finalized.is_set()) + + async def test_run_starts_session_with_bridge_identity(self) -> None: + runtime, _agent, executor = self._make_runtime() + runtime._ctx.room.remote_participants = {"user-1": object()} + + await runtime.run() + await asyncio.sleep(0) + await self._cancel_pending_tasks() + + start_commands = [cmd for cmd in executor.commands if isinstance(cmd, StartSession)] + self.assertEqual(len(start_commands), 1) + room_options = start_commands[0].room_options + participant_identity = getattr(room_options, "participant_identity", None) + if participant_identity is None: + participant_identity = room_options.kwargs["participant_identity"] + self.assertEqual(participant_identity, "bridge-1") + self.assertIn("user_input_transcribed", runtime._session.handlers) + + +if __name__ == "__main__": + unittest.main() + +#só pra fazer um novo PR diff --git a/tests/livekit/test_vad_flow_logging.py b/tests/livekit/test_vad_flow_logging.py new file mode 100644 index 0000000..44031a5 --- /dev/null +++ b/tests/livekit/test_vad_flow_logging.py @@ -0,0 +1,125 @@ +from __future__ import annotations + +import logging +from types import SimpleNamespace + +from livekit.agents import vad as agents_vad + +from app.livekit.main import FlowLoggingVADStream + + +def _event( + event_type: agents_vad.VADEventType, + *, + speech_s: float, + raw_speech_s: float | None = None, + silence_s: float = 0.0, + speaking: bool = False, +) -> SimpleNamespace: + return SimpleNamespace( + type=event_type, + speech_duration=speech_s, + raw_accumulated_speech=speech_s if raw_speech_s is None else raw_speech_s, + silence_duration=silence_s, + raw_accumulated_silence=silence_s, + probability=0.0, + speaking=speaking, + ) + + +def _stream(on_speech_end) -> FlowLoggingVADStream: + return FlowLoggingVADStream( + SimpleNamespace(), + stream_id=0, + should_log=True, + release_logging_stream=lambda _stream_id: None, + call_logger=logging.getLogger(__name__), + min_interrupt_s=0.5, + vad_config={ + "min_speech_duration": 0.15, + "min_silence_duration": 1.0, + "activation_threshold": 0.3, + "deactivation_threshold": 0.15, + }, + log_decisions=False, + log_activity=False, + activity_min_prob=0.0, + on_speech_end=on_speech_end, + ) + + +def test_speech_end_uses_final_duration_without_terminal_vad_silence() -> None: + durations_ms: list[int] = [] + stream = _stream(durations_ms.append) + + stream._log_vad_decision( + _event(agents_vad.VADEventType.START_OF_SPEECH, speech_s=0.15) + ) + # While Silero waits for the endpoint, INFERENCE_DONE keeps increasing + # speech_duration with the terminal silence. + stream._log_vad_decision( + _event( + agents_vad.VADEventType.INFERENCE_DONE, + speech_s=1.2, + raw_speech_s=0.2, + silence_s=1.0, + ) + ) + # END_OF_SPEECH reports the corrected duration after removing that silence. + stream._log_vad_decision( + _event( + agents_vad.VADEventType.END_OF_SPEECH, + speech_s=0.2, + raw_speech_s=0.2, + silence_s=1.0, + ) + ) + + assert durations_ms == [200] + + +def test_short_audio_stays_short_after_a_previous_long_round() -> None: + durations_ms: list[int] = [] + stream = _stream(durations_ms.append) + + stream._log_vad_decision( + _event(agents_vad.VADEventType.START_OF_SPEECH, speech_s=0.15) + ) + stream._log_vad_decision( + _event( + agents_vad.VADEventType.INFERENCE_DONE, + speech_s=2.2, + raw_speech_s=1.2, + silence_s=1.0, + ) + ) + stream._log_vad_decision( + _event( + agents_vad.VADEventType.END_OF_SPEECH, + speech_s=1.2, + raw_speech_s=1.2, + silence_s=1.0, + ) + ) + + stream._log_vad_decision( + _event(agents_vad.VADEventType.START_OF_SPEECH, speech_s=0.15) + ) + stream._log_vad_decision( + _event( + agents_vad.VADEventType.INFERENCE_DONE, + speech_s=1.2, + raw_speech_s=0.2, + silence_s=1.0, + ) + ) + stream._log_vad_decision( + _event( + agents_vad.VADEventType.END_OF_SPEECH, + speech_s=0.2, + raw_speech_s=0.2, + silence_s=1.0, + ) + ) + + assert durations_ms == [1200, 200] diff --git a/tests/livekit/test_wav_audio.py b/tests/livekit/test_wav_audio.py new file mode 100644 index 0000000..93d38fc --- /dev/null +++ b/tests/livekit/test_wav_audio.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +import asyncio +import sys +import types +import wave +from pathlib import Path +from unittest import mock + +from app.livekit.runtime.wav_audio import wav_audio_frames, wav_duration_ms + + +class _AudioFrame: + def __init__( + self, + *, + data: bytes, + sample_rate: int, + num_channels: int, + samples_per_channel: int, + ) -> None: + self.data = data + self.sample_rate = sample_rate + self.num_channels = num_channels + self.samples_per_channel = samples_per_channel + + +def _write_wav(path: Path, *, samples: int, sample_rate: int = 1000) -> None: + with wave.open(str(path), "wb") as wav: + wav.setnchannels(1) + wav.setsampwidth(2) + wav.setframerate(sample_rate) + wav.writeframes(b"\x01\x02" * samples) + + +def test_wav_audio_frames_pads_last_frame_and_adds_tail_silence(tmp_path: Path) -> None: + wav_path = tmp_path / "audio.wav" + _write_wav(wav_path, samples=25) + + livekit_module = types.ModuleType("livekit") + rtc_module = types.ModuleType("livekit.rtc") + rtc_module.AudioFrame = _AudioFrame + livekit_module.rtc = rtc_module + + async def _collect(): + with mock.patch.dict(sys.modules, {"livekit": livekit_module, "livekit.rtc": rtc_module}): + return [ + frame + async for frame in wav_audio_frames( + str(wav_path), + frame_duration_ms=20, + tail_silence_ms=40, + ) + ] + + frames = asyncio.run(_collect()) + + assert len(frames) == 4 + assert [frame.samples_per_channel for frame in frames] == [20, 20, 20, 20] + assert frames[1].data[:10] == b"\x01\x02" * 5 + assert frames[1].data[10:] == b"\x00" * 30 + assert frames[2].data == b"\x00" * 40 + assert frames[3].data == b"\x00" * 40 + assert wav_duration_ms(str(wav_path)) == 25 diff --git a/tests/providers/__init__.py b/tests/providers/__init__.py new file mode 100644 index 0000000..9d48db4 --- /dev/null +++ b/tests/providers/__init__.py @@ -0,0 +1 @@ +from __future__ import annotations diff --git a/tests/providers/test_fake_stt.py b/tests/providers/test_fake_stt.py new file mode 100644 index 0000000..4ac07ba --- /dev/null +++ b/tests/providers/test_fake_stt.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +import os +from unittest import mock +import unittest + +from livekit import rtc + +from app.providers.stt_fake import FakeSTT + + +def _audio_frame(samples_per_channel: int) -> rtc.AudioFrame: + return rtc.AudioFrame( + data=b"\x00\x00" * samples_per_channel, + sample_rate=16000, + num_channels=1, + samples_per_channel=samples_per_channel, + ) + + +class FakeSTTTests(unittest.IsolatedAsyncioTestCase): + async def test_fake_stt_returns_configured_transcripts_in_order(self) -> None: + with mock.patch.dict( + os.environ, + { + "FAKE_STT_TRANSCRIPTS": "primeira|segunda", + "FAKE_STT_MODE": "repeat_last", + "FAKE_STT_MIN_AUDIO_MS": "0", + }, + clear=True, + ): + stt = FakeSTT(language="pt-BR") + first = await stt.recognize([_audio_frame(3200)]) + second = await stt.recognize([_audio_frame(3200)]) + third = await stt.recognize([_audio_frame(3200)]) + + self.assertEqual(first.alternatives[0].text, "primeira") + self.assertEqual(second.alternatives[0].text, "segunda") + self.assertEqual(third.alternatives[0].text, "segunda") + + async def test_fake_stt_skips_too_short_audio(self) -> None: + with mock.patch.dict( + os.environ, + { + "FAKE_STT_TRANSCRIPTS": "fala", + "FAKE_STT_MIN_AUDIO_MS": "200", + }, + clear=True, + ): + stt = FakeSTT(language="pt-BR") + event = await stt.recognize([_audio_frame(800)]) + + self.assertEqual(event.alternatives[0].text, "") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/providers/test_internal_http_stt.py b/tests/providers/test_internal_http_stt.py new file mode 100644 index 0000000..f28972a --- /dev/null +++ b/tests/providers/test_internal_http_stt.py @@ -0,0 +1,465 @@ +from __future__ import annotations + +import json +import logging +import unittest +import uuid +from unittest import mock + +import httpx + +from app.providers.stt_internal_livekit import ( + InternalHTTPSTT, + InternalSTTConfig, + _prepend_pcm16le_silence, + stt_text_with_single_word_threshold, +) +from app.providers.stt_config import override_config +from app.utils import logging as logging_utils +from app.utils.turn_ids import ( + peek_started_turn_message_id, + register_started_turn_message_id, + reset_turn_message_sequence, +) + + +def _assert_uuid(testcase: unittest.TestCase, value: object) -> str: + text = str(value or "") + testcase.assertEqual(str(uuid.UUID(text)), text) + return text + + +def _single_word_payload(text: str, probability: float) -> dict: + return { + "data": { + "text": text, + "words": [{"word": text, "probability": probability}], + } + } + + +class InternalHTTPSTTTests(unittest.IsolatedAsyncioTestCase): + async def test_provider_can_be_instantiated(self) -> None: + client = httpx.AsyncClient() + try: + stt = InternalHTTPSTT( + InternalSTTConfig( + url="https://example.invalid/stt", + api_key="test-key", + ), + client=client, + ) + finally: + await client.aclose() + + self.assertEqual(stt.provider, "unknown") + + def test_default_sofya_vad_padding_keeps_more_prefix_audio(self) -> None: + vad_params = override_config["processor"]["config"]["extra_params"]["vad_parameters"] + + self.assertEqual(vad_params["speech_pad_ms"], 1000) + + def test_stt_input_prefix_padding_prepends_silence(self) -> None: + pcm = b"\x01\x00" * 16000 + + padded, padding_ms = _prepend_pcm16le_silence( + pcm, + sample_rate=16000, + channels=1, + padding_ms=250, + ) + + self.assertEqual(padding_ms, 250) + self.assertEqual(len(padded), len(pcm) + 8000) + self.assertTrue(padded.startswith(b"\x00" * 8000)) + self.assertTrue(padded.endswith(pcm)) + + async def test_config_override_is_used_as_json_object(self) -> None: + client = httpx.AsyncClient() + try: + stt = InternalHTTPSTT( + InternalSTTConfig( + url="https://example.invalid/stt", + api_key="test-key", + config_override='{"processor":{"strategy":"faster_default"}}', + ), + client=client, + ) + override = json.loads(stt._build_override_config_json()) + finally: + await client.aclose() + + self.assertEqual(override["processor"]["strategy"], "faster_default") + + async def test_http_success_with_empty_text_does_not_create_structured_turn(self) -> None: + published_events = [] + empty_transcript_calls = [] + captured_headers: dict[str, str] = {} + + def _handler(request: httpx.Request) -> httpx.Response: + captured_headers["connection"] = request.headers.get("connection", "") + return httpx.Response(200, json={"data": {"text": ""}}, request=request) + + logger = logging.getLogger("test.internal_http_stt.structured") + logger.handlers = [logging.NullHandler()] + logger.setLevel(logging.INFO) + structured_context = logging_utils.StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + reset_turn_message_sequence(structured_context, clear_pending=True) + upload_message_id = "12345678-1234-4234-9234-123456789abc" + register_started_turn_message_id(structured_context, upload_message_id) + + client = httpx.AsyncClient(transport=httpx.MockTransport(_handler)) + try: + stt = InternalHTTPSTT( + InternalSTTConfig( + url="https://example.invalid/stt", + api_key="test-key", + output_mode="api_text", + ), + client=client, + structured_log_context=structured_context, + structured_logger=logger, + empty_transcript_handler=lambda: empty_transcript_calls.append("called"), + ) + with ( + mock.patch("app.utils.logging.publish_structured_event", side_effect=published_events.append), + mock.patch("app.utils.logging.publish_structured_span"), + mock.patch("app.providers.stt_internal_livekit.log_flow_event") as flow_log, + mock.patch("app.providers.stt_internal_livekit.logger.info") as logger_info, + ): + event = await stt._post_internal_http( + b"\0\0", + b"fake-wav", + req_id="req-1", + started_ns=1_700_000_000_000_000_000, + language="pt-BR", + message_id=upload_message_id, + ) + finally: + await client.aclose() + + self.assertEqual(event.alternatives[0].text, "") + self.assertEqual(captured_headers["connection"], "close") + self.assertEqual(published_events, []) + self.assertEqual(peek_started_turn_message_id(structured_context), "") + self.assertTrue( + any( + call.args[1] == "stt_done" + and call.kwargs["request_id"] == "req-1" + and call.kwargs["text_len"] == 0 + for call in flow_log.call_args_list + ) + ) + self.assertEqual(empty_transcript_calls, ["called"]) + logger_info.assert_any_call( + "[stt][response_json] req_id=%s status=%s took=%.0fms json=%s", + "req-1", + 200, + mock.ANY, + '{"data": {"text": ""}}', + ) + + async def test_http_success_marks_structured_interruption(self) -> None: + published_events = [] + provider_metrics = [] + + def _handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"data": {"text": "sim"}}, request=request) + + logger = logging.getLogger("test.internal_http_stt.interruption") + logger.handlers = [logging.NullHandler()] + logger.setLevel(logging.INFO) + structured_context = logging_utils.StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + reset_turn_message_sequence(structured_context, clear_pending=True) + + client = httpx.AsyncClient(transport=httpx.MockTransport(_handler)) + try: + stt = InternalHTTPSTT( + InternalSTTConfig( + url="https://example.invalid/stt", + api_key="test-key", + output_mode="api_text", + ), + client=client, + structured_log_context=structured_context, + structured_logger=logger, + structured_interruption_flag=lambda: True, + metrics_handler=provider_metrics.append, + ) + with ( + mock.patch("app.utils.logging.publish_structured_event", side_effect=published_events.append), + mock.patch("app.utils.logging.publish_structured_span"), + ): + await stt._post_internal_http( + b"\0\0", + b"fake-wav", + req_id="req-1", + started_ns=1_700_000_000_000_000_000, + language="pt-BR", + audio_duration_ms=1250, + original_audio_duration_ms=1000, + input_padding_ms=250, + level_dbfs=-24.5, + ) + finally: + await client.aclose() + + self.assertEqual(len(published_events), 1) + self.assertEqual(published_events[0]["interrupcao"], 1) + self.assertEqual(len(provider_metrics), 1) + self.assertEqual(provider_metrics[0]["event"], "completed") + self.assertEqual(provider_metrics[0]["audio_duration_ms"], 1250) + self.assertEqual(provider_metrics[0]["original_audio_duration_ms"], 1000) + self.assertEqual(provider_metrics[0]["input_padding_ms"], 250) + self.assertEqual(provider_metrics[0]["input_dbfs"], -24.5) + self.assertEqual(provider_metrics[0]["retry_count"], 0) + self.assertEqual(provider_metrics[0]["http_status"], 200) + self.assertFalse(provider_metrics[0]["empty_transcript"]) + self.assertEqual(provider_metrics[0]["text_length"], 3) + + async def test_http_success_uses_supplied_message_id_for_structured_event(self) -> None: + published_events = [] + + def _handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"data": {"text": "sim"}}, request=request) + + logger = logging.getLogger("test.internal_http_stt.message_id") + logger.handlers = [logging.NullHandler()] + logger.setLevel(logging.INFO) + structured_context = logging_utils.StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + reset_turn_message_sequence(structured_context, clear_pending=True) + + client = httpx.AsyncClient(transport=httpx.MockTransport(_handler)) + try: + stt = InternalHTTPSTT( + InternalSTTConfig( + url="https://example.invalid/stt", + api_key="test-key", + output_mode="api_text", + ), + client=client, + structured_log_context=structured_context, + structured_logger=logger, + ) + with ( + mock.patch("app.utils.logging.publish_structured_event", side_effect=published_events.append), + mock.patch("app.utils.logging.publish_structured_span"), + ): + await stt._post_internal_http( + b"\0\0", + b"fake-wav", + req_id="req-1", + started_ns=1_700_000_000_000_000_000, + language="pt-BR", + message_id="message-from-upload-path", + ) + finally: + await client.aclose() + + self.assertEqual(len(published_events), 1) + self.assertEqual(published_events[0]["message_id"], "message-from-upload-path") + + async def test_http_500_retries_once_then_uses_success(self) -> None: + attempts = 0 + + def _handler(request: httpx.Request) -> httpx.Response: + nonlocal attempts + attempts += 1 + if attempts == 1: + return httpx.Response(500, text="temporary", request=request) + return httpx.Response(200, json={"data": {"text": "sim"}}, request=request) + + client = httpx.AsyncClient(transport=httpx.MockTransport(_handler)) + try: + stt = InternalHTTPSTT( + InternalSTTConfig( + url="https://example.invalid/stt", + api_key="test-key", + output_mode="api_text", + ), + client=client, + ) + with mock.patch("app.providers.stt_internal_livekit.log_flow_event") as flow_log: + event = await stt._post_internal_http( + b"\0\0", + b"fake-wav", + req_id="req-retry", + started_ns=1_700_000_000_000_000_000, + language="pt-BR", + ) + finally: + await client.aclose() + + self.assertEqual(attempts, 2) + self.assertEqual(event.alternatives[0].text, "sim") + self.assertTrue( + any( + call.args[1] == "stt_http_retry" + and call.kwargs["request_id"] == "req-retry" + and call.kwargs["attempt"] == 1 + and call.kwargs["max_retries"] == 1 + for call in flow_log.call_args_list + ) + ) + + async def test_http_500_after_retry_returns_empty_transcript_without_raising(self) -> None: + attempts = 0 + published_events = [] + + def _handler(request: httpx.Request) -> httpx.Response: + nonlocal attempts + attempts += 1 + return httpx.Response(500, text="still failing", request=request) + + logger = logging.getLogger("test.internal_http_stt.http_500_nonfatal") + logger.handlers = [logging.NullHandler()] + logger.setLevel(logging.INFO) + structured_context = logging_utils.StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ) + reset_turn_message_sequence(structured_context, clear_pending=True) + + client = httpx.AsyncClient(transport=httpx.MockTransport(_handler)) + try: + stt = InternalHTTPSTT( + InternalSTTConfig( + url="https://example.invalid/stt", + api_key="test-key", + output_mode="api_text", + ), + client=client, + structured_log_context=structured_context, + structured_logger=logger, + ) + with ( + mock.patch("app.utils.logging.publish_structured_event", side_effect=published_events.append), + mock.patch("app.utils.logging.publish_structured_span"), + mock.patch("app.providers.stt_internal_livekit.log_flow_event") as flow_log, + ): + event = await stt._post_internal_http( + b"\0\0", + b"fake-wav", + req_id="req-500", + started_ns=1_700_000_000_000_000_000, + language="pt-BR", + ) + finally: + await client.aclose() + + self.assertEqual(attempts, 2) + self.assertEqual(event.alternatives[0].text, "") + self.assertEqual(len(published_events), 1) + self.assertEqual(published_events[0]["erro_msg"], "Falha STT") + self.assertEqual(published_events[0]["http_cod_status"], 500) + self.assertTrue( + any( + call.args[1] == "stt_error_nonfatal" + and call.kwargs["request_id"] == "req-500" + and call.kwargs["action"] == "return_empty_transcript" + for call in flow_log.call_args_list + ) + ) + + def test_single_word_allowlist_accepts_low_confidence_sim(self) -> None: + self.assertEqual( + stt_text_with_single_word_threshold( + _single_word_payload("sim", 0.02), + min_prob_single_word=0.03, + ), + "sim", + ) + + def test_single_word_allowlist_accepts_extremely_low_confidence_sim(self) -> None: + self.assertEqual( + stt_text_with_single_word_threshold( + _single_word_payload("sim", 0.009), + min_prob_single_word=0.03, + ), + "sim", + ) + + def test_single_word_filter_rejects_non_allowlisted_low_confidence_word(self) -> None: + self.assertIsNone( + stt_text_with_single_word_threshold( + _single_word_payload("talvez", 0.02), + min_prob_single_word=0.03, + ) + ) + + def test_single_word_filter_accepts_non_allowlisted_high_confidence_word(self) -> None: + self.assertEqual( + stt_text_with_single_word_threshold( + _single_word_payload("talvez", 0.04), + min_prob_single_word=0.03, + ), + "talvez", + ) + + def test_single_word_filter_keeps_api_text_empty_as_empty(self) -> None: + self.assertIsNone( + stt_text_with_single_word_threshold( + { + "data": { + "text": "", + "words": [{"word": "sim", "probability": 0.5}], + } + }, + min_prob_single_word=0.03, + ) + ) + + async def test_single_word_allowlist_logs_low_confidence_reason(self) -> None: + client = httpx.AsyncClient() + try: + stt = InternalHTTPSTT( + InternalSTTConfig( + url="https://example.invalid/stt", + api_key="test-key", + min_prob_single_word=0.03, + ), + client=client, + ) + with mock.patch("app.providers.stt_internal_livekit.log_flow_event") as flow_log: + text = stt._format_stt_output( + _single_word_payload("sim", 0.02), + request_id="req-allowlist", + ) + finally: + await client.aclose() + + self.assertEqual(text, "sim") + self.assertTrue( + any( + call.args[1] == "stt_payload" + and call.kwargs["request_id"] == "req-allowlist" + and call.kwargs["filter_reason"] == "single_word_allowlist_low_confidence" + and call.kwargs["allowlist_min_prob"] == 0.0 + for call in flow_log.call_args_list + ) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/providers/test_tts.py b/tests/providers/test_tts.py new file mode 100644 index 0000000..d88467b --- /dev/null +++ b/tests/providers/test_tts.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +import os +from unittest import mock +import unittest + +from app.providers import tts as tts_module +from app.providers.tts import FakeTTS, build_tts_provider_from_env + + +class ProviderTTSTests(unittest.TestCase): + def test_build_tts_provider_from_env_returns_reason_when_provider_is_unsupported(self) -> None: + with mock.patch.dict(os.environ, {}, clear=True): + provider, reason = build_tts_provider_from_env("azure") + + self.assertIsNone(provider) + self.assertEqual(reason, "unsupported_tts_provider:azure") + + def test_build_tts_provider_from_env_returns_reason_when_elevenlabs_sdk_is_missing(self) -> None: + with mock.patch.dict( + os.environ, + { + "ELEVENLABS_API_KEY": "key-123", + "ELEVENLABS_VOICE_ID": "voice-123", + }, + clear=True, + ): + with mock.patch.object(tts_module, "_is_elevenlabs_available", return_value=False): + provider, reason = build_tts_provider_from_env("elevenlabs") + + self.assertIsNone(provider) + self.assertEqual(reason, "missing_elevenlabs_sdk") + + def test_build_tts_provider_from_env_returns_fake_provider(self) -> None: + with mock.patch.dict( + os.environ, + { + "FAKE_TTS_TONE_HZ": "512", + "FAKE_TTS_CHAR_DURATION_MS": "18", + }, + clear=True, + ): + provider, reason = build_tts_provider_from_env("fake") + + self.assertIsInstance(provider, FakeTTS) + self.assertIsNone(reason) + self.assertGreater(len(provider.synthesize_pcm16k("teste fake")), 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/services/test_session_context.py b/tests/services/test_session_context.py new file mode 100644 index 0000000..87abb3f --- /dev/null +++ b/tests/services/test_session_context.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +from app.services.session_context import extract_protocol + + +def test_extract_protocol_prefers_start_payload_data() -> None: + protocol = extract_protocol( + { + "data": { + "protocolo": "PRT-123", + } + }, + {"protocolo": "PRT-999"}, + ) + + assert protocol == "PRT-123" + + +def test_extract_protocol_falls_back_to_session_data() -> None: + protocol = extract_protocol( + {"data": {}}, + {"protocolo": "PRT-999"}, + ) + + assert protocol == "PRT-999" + + +def test_extract_protocol_falls_back_to_router_call_key_when_protocol_is_missing() -> None: + protocol = extract_protocol( + { + "data": { + "routerCallKey": "RCK-123", + } + }, + {}, + ) + + assert protocol == "RCK-123" + + +def test_extract_protocol_accepts_protocol_id() -> None: + protocol = extract_protocol( + { + "data": { + "protocol_id": "PRT-777", + "routerCallKey": "RCK-123", + } + }, + {}, + ) + + assert protocol == "PRT-777" diff --git a/tests/tools/__init__.py b/tests/tools/__init__.py new file mode 100644 index 0000000..4dbf6a8 --- /dev/null +++ b/tests/tools/__init__.py @@ -0,0 +1 @@ +"""Tests for operational tooling.""" diff --git a/tests/tools/test_local_stresstest.py b/tests/tools/test_local_stresstest.py new file mode 100644 index 0000000..6418317 --- /dev/null +++ b/tests/tools/test_local_stresstest.py @@ -0,0 +1,258 @@ +from __future__ import annotations + +import asyncio +import json +import math +import wave +from pathlib import Path + +from app.tools.local_stresstest.audio import ( + AudioSample, + audio_metrics, + build_variations, + read_wav_mono16, + vad_proxy_metrics, + write_wav, +) +from app.tools.local_stresstest.report import render_markdown_report, write_csv, write_mermaid_files +from app.tools.local_stresstest.runner import StressConfig, wait_for_local_services +from app.tools.local_stresstest.scenarios import scenario_description +from app.tools.local_stresstest.text import compare_text, normalize_text, word_error_rate +from app.tools.local_stresstest.timeline import excerpt_timeline, read_timeline, timeline_has_error + + +def _tone_pcm(*, sample_rate: int = 16_000, duration_ms: int = 300, hz: float = 440.0) -> bytes: + samples = round(sample_rate * duration_ms / 1000) + out = bytearray() + for idx in range(samples): + value = round(math.sin(2 * math.pi * hz * idx / sample_rate) * 9000) + out.extend(int(value).to_bytes(2, byteorder="little", signed=True)) + return bytes(out) + + +def test_normalize_text_removes_accents_case_and_punctuation() -> None: + assert normalize_text(" Teste UNITÁRIO, Sofya! ") == "teste unitario sofya" + + +def test_word_error_rate_counts_insertions_deletions_and_substitutions() -> None: + wer, substitutions, deletions, insertions = word_error_rate( + "teste unitario do stt sofya", + "teste unitario stt sofia agora", + ) + + assert round(wer, 2) == 0.6 + assert substitutions == 1 + assert deletions == 1 + assert insertions == 1 + + +def test_compare_text_reports_missing_critical_terms() -> None: + comparison = compare_text( + expected="Teste unitario do STT Sofya", + actual="Teste unitario do STT", + critical_terms=["teste", "sofya"], + ) + + assert comparison.terms_ok is False + assert comparison.missing_terms == ("sofya",) + + +def test_audio_variations_keep_expected_names_and_are_nonempty() -> None: + sample = AudioSample(name="base", pcm=_tone_pcm()) + variations = build_variations(sample) + + assert [item.name for item in variations] == [ + "clean", + "low_volume", + "very_low_volume", + "high_volume", + "clipped_high_volume", + "leading_trailing_silence", + "short_leading_silence", + "long_leading_silence", + "noise_snr_20", + "noise_snr_15", + "noise_snr_10", + "pre_noise_300ms", + "pre_noise_800ms", + "telephony_profile", + "telephony_low_volume", + "initial_fade_in_250ms", + "initial_fade_in_500ms", + "initial_fade_in_900ms", + "initial_dip_300ms", + "initial_dip_700ms", + "prefix_300ms_15pct", + "prefix_600ms_20pct", + "prefix_900ms_25pct", + "low_prefix_noise_600ms", + "low_prefix_telephony_700ms", + ] + assert all(item.pcm for item in variations) + assert len(variations[5].pcm) > len(sample.pcm) + assert len(variations) == 25 + + +def test_audio_metrics_detects_audible_tone_and_duration() -> None: + metrics = audio_metrics(_tone_pcm(duration_ms=500)) + + assert metrics.duration_ms == 500 + assert metrics.audible is True + assert metrics.clipping_ratio == 0.0 + + +def test_vad_proxy_metrics_flags_soft_start_risk() -> None: + sample = AudioSample(name="base", pcm=(b"\x00" * 1600) + _tone_pcm(duration_ms=300)) + metrics = vad_proxy_metrics(sample, threshold_dbfs=-45.0, prefix_padding_ms=20) + + assert metrics.first_voice_ms > 0 + assert metrics.unrecovered_prefix_ms > 0 + assert metrics.low_start_risk is True + + +def test_wav_read_converts_to_16k_mono(tmp_path: Path) -> None: + stereo_path = tmp_path / "stereo.wav" + left = _tone_pcm(sample_rate=8_000, duration_ms=100) + stereo = bytearray() + for idx in range(0, len(left), 2): + stereo.extend(left[idx : idx + 2]) + stereo.extend(left[idx : idx + 2]) + with wave.open(str(stereo_path), "wb") as handle: + handle.setnchannels(2) + handle.setsampwidth(2) + handle.setframerate(8_000) + handle.writeframes(bytes(stereo)) + + sample = read_wav_mono16(stereo_path) + + assert sample.sample_rate == 16_000 + assert sample.channels == 1 + assert audio_metrics(sample.pcm).duration_ms == 100 + + +def test_write_wav_creates_parent_directory(tmp_path: Path) -> None: + out = write_wav(tmp_path / "nested" / "tone.wav", AudioSample(name="tone", pcm=_tone_pcm())) + + assert out.exists() + assert read_wav_mono16(out).pcm + + +def test_render_markdown_report_contains_tables_and_mermaid() -> None: + report = render_markdown_report( + { + "passed": True, + "started_at": "2026-06-16T00:00:00Z", + "duration_ms": 123, + "expected_text": "Teste", + "baseline": {"path": "baseline.wav", "synthetic": True}, + "stt_results": [{"scenario": "clean", "passed": True, "wer": 0.0}], + "tts_results": [{"scenario": "short", "passed": True, "duration_ms": 1000}], + "e2e_results": [{"scenario": "clean", "passed": True, "ready_received": True}], + "startup_checks": { + "bridge": { + "service": "bridge", + "url": "http://127.0.0.1:8000/health", + "status": "ready", + "attempts": 1, + "detail": "HTTP 200", + } + }, + "artifacts": {"summary_json": "summary.json"}, + } + ) + + assert "Aviso: esta execucao usou um audio base sintetico" in report + assert "## Prontidao Local" in report + assert "Versao textual:" in report + assert "![Fluxo STT](diagrams/stt_flow.svg)" in report + assert "Codigo Mermaid:" in report + assert "```mermaid" in report + assert "| cenario | descricao | status | prefixo_ok |" in report + assert "Audio base sem degradacao" in report + + +def test_scenario_description_ignores_repeat_suffix() -> None: + assert scenario_description("noise_snr_20_r2").startswith("Audio com ruido leve") + + +def test_write_mermaid_files_creates_standalone_diagrams(tmp_path: Path) -> None: + files = write_mermaid_files(tmp_path) + + assert Path(files["stt_mermaid"]).read_text(encoding="utf-8").startswith("flowchart LR") + assert "sequenceDiagram" in Path(files["e2e_mermaid"]).read_text(encoding="utf-8") + assert Path(files["stt_svg"]).read_text(encoding="utf-8").startswith(" StressConfig: + return StressConfig( + env_file=Path(".env.dev"), + report_dir=tmp_path, + expected_text="Teste", + synthesis_text="Teste", + critical_terms=("teste",), + stt_wer_threshold=0.2, + bridge_url="ws://127.0.0.1:8000/ws/agent", + bridge_health_url="http://bridge.local/health", + agent_health_url="http://agent.local/", + startup_wait_s=5.0, + startup_poll_s=0.01, + skip_local_wait=False, + repeat=1, + concurrency=1, + e2e_turns=1, + e2e_timeout_s=5.0, + stress_audio=None, + prefix_text="Teste", + prefix_words=1, + vad_proxy_threshold_dbfs=-45.0, + vad_proxy_prefix_padding_ms=1000, + vad_proxy_min_speech_ms=100, + ) + + +def test_wait_for_local_services_retries_until_agent_is_ready(tmp_path: Path) -> None: + config = _stress_config(tmp_path) + attempts: dict[str, int] = {} + + async def probe(url: str) -> tuple[bool, str]: + attempts[url] = attempts.get(url, 0) + 1 + if "agent" in url and attempts[url] < 3: + return False, "connection refused" + return True, "HTTP 200" + + async def no_sleep(_: float) -> None: + return None + + states = asyncio.run(wait_for_local_services(config, probe=probe, sleep=no_sleep)) + + assert states["bridge"]["ok"] is True + assert states["bridge"]["attempts"] == 1 + assert states["agent_runtime"]["ok"] is True + assert states["agent_runtime"]["attempts"] == 3 + + +def test_write_csv_handles_empty_rows(tmp_path: Path) -> None: + out = write_csv(tmp_path / "empty.csv", []) + + assert out.read_text(encoding="utf-8").strip() == "" + + +def test_timeline_helpers_parse_excerpt_and_errors(tmp_path: Path) -> None: + path = tmp_path / "timeline.jsonl" + records = [ + {"event": "ready_sent"}, + {"event": "noise"}, + {"event": "user_transcript_final", "text": "ola"}, + {"event": "bridge_failed"}, + ] + path.write_text("\n".join(json.dumps(item) for item in records), encoding="utf-8") + + parsed = read_timeline(path) + + assert parsed == records + assert [item["event"] for item in excerpt_timeline(parsed)] == [ + "ready_sent", + "user_transcript_final", + "bridge_failed", + ] + assert timeline_has_error(parsed) is True diff --git a/tests/tools/test_oci_audio_download.py b/tests/tools/test_oci_audio_download.py new file mode 100644 index 0000000..96799b4 --- /dev/null +++ b/tests/tools/test_oci_audio_download.py @@ -0,0 +1,119 @@ +from __future__ import annotations + +import zipfile +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from app.tools.oci_audio_download import ( + _parse_date, + _parse_session_id, + download_prefix, + entire_calls_prefix, + segments_prefix, +) +from app.utils.stt_audio_upload import OCIUploadConfig + + +class _Raw: + def __init__(self, content: bytes) -> None: + self.content = content + + def stream(self, _size: int, decode_content: bool = False): + assert decode_content is False + yield self.content + + +class _Client: + def __init__(self, objects: dict[str, bytes]) -> None: + self.objects = objects + self.downloaded: list[str] = [] + + def list_objects(self, **kwargs): + prefix = kwargs["prefix"] + items = [ + SimpleNamespace(name=name, size=len(content), etag="etag") + for name, content in self.objects.items() + if name.startswith(prefix) + ] + return SimpleNamespace(data=SimpleNamespace(objects=items, next_start_with=None)) + + def get_object(self, **kwargs): + name = kwargs["object_name"] + self.downloaded.append(name) + return SimpleNamespace(data=SimpleNamespace(raw=_Raw(self.objects[name]))) + + +def _config() -> OCIUploadConfig: + return OCIUploadConfig("local", "sa-saopaulo-1", "tia-audio", "namespace") + + +def test_prefixes_match_object_storage_layout() -> None: + assert segments_prefix("2026-07-24", "session-123") == "2026-07-24/session-123/" + assert entire_calls_prefix("2026-07-24") == "2026-07-24/entire_call/" + + +def test_argument_validation() -> None: + assert _parse_date("2026-07-24") == "2026-07-24" + assert _parse_session_id("session_123") == "session_123" + with pytest.raises(Exception): + _parse_date("24/07/2026") + with pytest.raises(Exception): + _parse_session_id("../entire_call") + + +def test_download_prefix_downloads_and_creates_zip(tmp_path: Path) -> None: + prefix = "2026-07-24/entire_call/" + client = _Client( + { + f"{prefix}call-a.wav": b"audio-a", + f"{prefix}call-b.wav": b"audio-b", + "2026-07-23/entire_call/old.wav": b"old", + } + ) + output_dir = tmp_path / "calls" + zip_path = tmp_path / "calls.zip" + result = download_prefix( + client=client, + config=_config(), + prefix=prefix, + output_dir=output_dir, + zip_path=zip_path, + ) + assert result.object_count == 2 + assert result.downloaded_count == 2 + assert (output_dir / "call-a.wav").read_bytes() == b"audio-a" + with zipfile.ZipFile(zip_path) as archive: + assert sorted(archive.namelist()) == ["call-a.wav", "call-b.wav"] + assert archive.testzip() is None + + +def test_download_prefix_reuses_complete_file(tmp_path: Path) -> None: + prefix = "2026-07-24/session-123/" + name = f"{prefix}message.wav" + client = _Client({name: b"segment"}) + output_dir = tmp_path / "segments" + output_dir.mkdir() + (output_dir / "message.wav").write_bytes(b"segment") + result = download_prefix( + client=client, + config=_config(), + prefix=prefix, + output_dir=output_dir, + zip_path=None, + ) + assert result.cached_count == 1 + assert result.downloaded_count == 0 + assert client.downloaded == [] + + +def test_download_prefix_rejects_empty_prefix(tmp_path: Path) -> None: + with pytest.raises(FileNotFoundError): + download_prefix( + client=_Client({}), + config=_config(), + prefix="2026-07-24/entire_call/", + output_dir=tmp_path, + zip_path=None, + ) diff --git a/tests/utils/__init__.py b/tests/utils/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/utils/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/utils/test_audio_backlog.py b/tests/utils/test_audio_backlog.py new file mode 100644 index 0000000..18230d4 --- /dev/null +++ b/tests/utils/test_audio_backlog.py @@ -0,0 +1,109 @@ +from __future__ import annotations + +import asyncio +import struct + +from app.utils import audio_backlog + +FRAME_SAMPLES = 320 # 20 ms @ 16 kHz mono +BYTES_PER_FRAME = FRAME_SAMPLES * 2 + + +def _silent() -> bytes: + return b"\x00" * BYTES_PER_FRAME + + +def _voiced(amp: int = 8000) -> bytes: + return struct.pack("<%dh" % FRAME_SAMPLES, *([amp] * FRAME_SAMPLES)) + + +def _queue(frames: list[bytes]) -> asyncio.Queue: + q: asyncio.Queue = asyncio.Queue() + for frame in frames: + q.put_nowait(frame) + return q + + +def _drain(q: asyncio.Queue) -> list[bytes]: + out = [] + while True: + try: + out.append(q.get_nowait()) + except asyncio.QueueEmpty: + break + return out + + +def test_rms_threshold_from_dbfs() -> None: + assert audio_backlog.rms_threshold_from_dbfs(0.0) == 32768 + assert audio_backlog.rms_threshold_from_dbfs(-50.0) == 103 + assert audio_backlog.rms_threshold_from_dbfs(-200.0) == 0 + + +def test_frames_from_ms_rounds_up_with_minimum() -> None: + assert audio_backlog.frames_from_ms(100, 20) == 5 + assert audio_backlog.frames_from_ms(90, 20) == 5 # arredonda pra cima + assert audio_backlog.frames_from_ms(0, 20) == 1 # minimo + assert audio_backlog.frames_from_ms(40, 20) == 2 + + +def test_is_silent_frame() -> None: + thr = audio_backlog.rms_threshold_from_dbfs(-50.0) + assert audio_backlog.is_silent_frame(_silent(), thr) is True + assert audio_backlog.is_silent_frame(_voiced(), thr) is False + + +def test_shed_below_keep_is_noop() -> None: + q = _queue([_voiced()] * 2) + result = audio_backlog.shed_queue_backlog(q, keep_frames=5, rms_threshold=103, blind=False) + assert result.dropped == 0 + assert result.mode == "none" + assert q.qsize() == 2 + + +def test_shed_blind_drops_oldest_and_keeps_recent_in_order() -> None: + frames = [_voiced(1000 + i) for i in range(8)] + q = _queue(frames) + result = audio_backlog.shed_queue_backlog(q, keep_frames=2, rms_threshold=103, blind=True) + assert result.dropped == 6 + assert result.dropped_voiced == 6 + assert result.mode == "blind" + assert _drain(q) == frames[-2:] + + +def test_shed_energy_drops_only_silence_preserving_speech_order() -> None: + voiced = [_voiced(2000 + i * 100) for i in range(4)] + # intercala voz e silencio: [V0, S, V1, S, V2, S, V3, S] + frames = [voiced[0], _silent(), voiced[1], _silent(), voiced[2], _silent(), voiced[3], _silent()] + q = _queue(frames) + result = audio_backlog.shed_queue_backlog(q, keep_frames=2, rms_threshold=103, blind=False) + # need = 8 - 2 = 6, mas so ha 4 frames silenciosos -> descarta os 4 silencios, + # preserva os 4 de voz na ordem (nao corta fala). + assert result.mode == "energy" + assert result.dropped == 4 + assert result.dropped_silent == 4 + assert result.dropped_voiced == 0 + assert _drain(q) == voiced + + +def test_shed_energy_stops_at_need_even_with_more_silence() -> None: + voiced = [_voiced(3000 + i) for i in range(2)] + # 2 de voz + 6 de silencio, keep=4 -> need=4 -> descarta 4 silencios, sobra 4 + frames = [voiced[0], _silent(), _silent(), _silent(), voiced[1], _silent(), _silent(), _silent()] + q = _queue(frames) + result = audio_backlog.shed_queue_backlog(q, keep_frames=4, rms_threshold=103, blind=False) + assert result.dropped == 4 + assert result.dropped_silent == 4 + remaining = _drain(q) + assert q.qsize() == 0 + # sobra os 2 de voz (na ordem) + 2 silencios mais recentes + assert remaining[0] == voiced[0] + assert voiced[1] in remaining + assert len(remaining) == 4 + + +def test_shed_disabled_via_empty_queue() -> None: + q = _queue([]) + result = audio_backlog.shed_queue_backlog(q, keep_frames=2, rms_threshold=103, blind=True) + assert result.dropped == 0 + assert result.mode == "none" diff --git a/tests/utils/test_audio_output_tracker.py b/tests/utils/test_audio_output_tracker.py new file mode 100644 index 0000000..bbd49f7 --- /dev/null +++ b/tests/utils/test_audio_output_tracker.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +import logging + +from app.utils import background as background_module +from app.utils.background import ( + AudioOutputBacklogConfig, + AudioOutputLatencyTracker, + audio_output_backlog_config_from_env, +) + +_LOGGER = logging.getLogger("test.audio_out") + + +def _capture(monkeypatch) -> list: + events: list = [] + monkeypatch.setattr( + background_module, + "log_flow_event", + lambda logger, step, **payload: events.append((step, payload)), + ) + return events + + +def test_output_config_from_env_defaults(monkeypatch) -> None: + for name in ( + "AUDIO_OUT_LATENCY_METRICS_ENABLED", + "AUDIO_OUT_LATENCY_ALERT_MS", + "AUDIO_OUT_LATENCY_LOG_INTERVAL_S", + "AUDIO_OUT_PACING_DEBT_THRESHOLD_MS", + "AUDIO_OUT_BACKLOG_SHED_ENABLED", + "AUDIO_OUT_BACKLOG_SHED_THRESHOLD_MS", + "AUDIO_OUT_BACKLOG_SHED_KEEP_MS", + "AUDIO_OUT_BACKLOG_SHED_CHECK_INTERVAL_FRAMES", + "AUDIO_OUT_BACKLOG_SILENCE_DBFS", + ): + monkeypatch.delenv(name, raising=False) + + config = audio_output_backlog_config_from_env() + + assert config.metrics_enabled is True + assert config.shed_enabled is True + assert config.shed_threshold_ms == 400 + assert config.shed_keep_ms == 200 + assert config.pacing_debt_threshold_ms == 200 + assert config.silence_dbfs == -60.0 + # frames dependem do frame_ms passado no uso + assert config.shed_threshold_frames(20) == 20 + assert config.shed_keep_frames(20) == 10 + assert config.silence_rms_threshold == 32 + + +def test_note_pacing_debt_accumulates_and_logs(monkeypatch) -> None: + events = _capture(monkeypatch) + tracker = AudioOutputLatencyTracker( + AudioOutputBacklogConfig(), frame_ms=20, flow_logger=_LOGGER + ) + + tracker.note_pacing_debt(250, now=1.0) + tracker.note_pacing_debt(300, now=2.0) + tracker.note_pacing_debt(0, now=3.0) # ignora nao-positivo + + assert tracker.pacing_debt_events == 2 + assert tracker.pacing_debt_dropped_ms == 550 + debt_events = [e for e in events if e[0] == "audio_out_pacing_debt"] + assert len(debt_events) == 2 + assert debt_events[-1][1]["pacing_debt_dropped_ms"] == 550 + assert debt_events[-1][1]["debt_ms"] == 300 + + +def test_record_shed_accumulates_and_logs(monkeypatch) -> None: + events = _capture(monkeypatch) + debug_events = [] + tracker = AudioOutputLatencyTracker( + AudioOutputBacklogConfig(), + frame_ms=20, + flow_logger=_LOGGER, + debug_event_publisher=lambda event, payload: debug_events.append( + (event, payload) + ), + ) + + tracker.record_shed( + dropped_frames=3, + dropped_silent=3, + mode="energy", + queue_before_frames=30, + queue_after_frames=27, + ) + tracker.record_shed( + dropped_frames=0, # ignora + dropped_silent=0, + mode="energy", + queue_before_frames=10, + queue_after_frames=10, + ) + + assert tracker.total_dropped_frames == 3 + shed_events = [e for e in events if e[0] == "audio_out_latency_shed"] + assert len(shed_events) == 1 + payload = shed_events[0][1] + assert payload["mode"] == "energy" + assert payload["dropped_ms"] == 60 + assert payload["dropped_silent_frames"] == 3 + assert payload["queue_before_ms"] == 600 + assert payload["total_dropped_ms"] == 60 + assert debug_events[-1][0] == "bridge.audio_out.shed" + assert debug_events[-1][1]["dropped_ms"] == 60 + + +def test_maybe_log_primes_first_then_logs_on_interval(monkeypatch) -> None: + events = _capture(monkeypatch) + config = AudioOutputBacklogConfig(latency_alert_ms=1000, latency_log_interval_s=15.0) + tracker = AudioOutputLatencyTracker(config, frame_ms=20, flow_logger=_LOGGER) + + # primeira chamada apenas inicializa o relogio (sem log) + tracker.maybe_log(queue_frames=5, now=100.0, reason="tick") + assert [e for e in events if e[0] == "audio_out_latency"] == [] + assert tracker.peak_queue_ms == 100 # peak atualiza mesmo sem logar + + # antes do intervalo: ainda nao loga + tracker.maybe_log(queue_frames=5, now=110.0, reason="tick") + assert [e for e in events if e[0] == "audio_out_latency"] == [] + + # passado o intervalo: loga + tracker.maybe_log(queue_frames=5, now=116.0, reason="tick") + out_events = [e for e in events if e[0] == "audio_out_latency"] + assert len(out_events) == 1 + assert out_events[0][1]["queue_ms"] == 100 + + +def test_maybe_log_logs_on_alert_threshold(monkeypatch) -> None: + events = _capture(monkeypatch) + config = AudioOutputBacklogConfig(latency_alert_ms=1000, latency_log_interval_s=999.0) + tracker = AudioOutputLatencyTracker(config, frame_ms=20, flow_logger=_LOGGER) + + tracker.maybe_log(queue_frames=5, now=1.0, reason="tick") # prime + # backlog >= alert (60 frames * 20 = 1200 ms >= 1000) -> loga fora do intervalo + tracker.maybe_log(queue_frames=60, now=1.5, reason="tick") + + out_events = [e for e in events if e[0] == "audio_out_latency"] + assert len(out_events) == 1 + assert out_events[0][1]["queue_ms"] == 1200 + assert out_events[0][1]["peak_queue_ms"] == 1200 + + +def test_metrics_disabled_suppresses_all_logs(monkeypatch) -> None: + events = _capture(monkeypatch) + config = AudioOutputBacklogConfig(metrics_enabled=False) + tracker = AudioOutputLatencyTracker(config, frame_ms=20, flow_logger=_LOGGER) + + tracker.note_pacing_debt(500, now=1.0) + tracker.record_shed( + dropped_frames=5, dropped_silent=5, mode="energy", + queue_before_frames=40, queue_after_frames=35, + ) + tracker.maybe_log(queue_frames=100, now=1.0, reason="tick", force=True) + + assert events == [] + # contadores ainda acumulam (metrica desligada nao perde estado interno) + assert tracker.pacing_debt_dropped_ms == 500 + assert tracker.total_dropped_frames == 5 + + +def test_maybe_log_alert_is_rate_limited_until_it_recovers(monkeypatch) -> None: + events = _capture(monkeypatch) + config = AudioOutputBacklogConfig(latency_alert_ms=1000, latency_log_interval_s=15.0) + tracker = AudioOutputLatencyTracker(config, frame_ms=20, flow_logger=_LOGGER) + + tracker.maybe_log(queue_frames=5, now=1.0, reason="tick") + tracker.maybe_log(queue_frames=60, now=1.5, reason="alert_start") + tracker.maybe_log(queue_frames=60, now=1.52, reason="alert_still_active") + tracker.maybe_log(queue_frames=5, now=2.0, reason="recovered") + tracker.maybe_log(queue_frames=60, now=2.1, reason="alert_restart") + + out_events = [event for event in events if event[0] == "audio_out_latency"] + assert len(out_events) == 2 + assert [event[1]["reason"] for event in out_events] == ["alert_start", "alert_restart"] diff --git a/tests/utils/test_background.py b/tests/utils/test_background.py new file mode 100644 index 0000000..21ce6a7 --- /dev/null +++ b/tests/utils/test_background.py @@ -0,0 +1,152 @@ +from __future__ import annotations + +import asyncio +import struct + +from fastapi import WebSocketDisconnect + +from app.utils.background import BridgeOutputStats, ws_out_loop + + +class _ClosingWebSocket: + def __init__(self, exc: BaseException) -> None: + self.exc = exc + self.sent_frames = 0 + + async def send_bytes(self, frame: bytes) -> None: + self.sent_frames += 1 + raise self.exc + + +class _CollectThenDisconnectWebSocket: + def __init__(self, disconnect_after: int) -> None: + self.disconnect_after = disconnect_after + self.sent_frames: list[bytes] = [] + + async def send_bytes(self, frame: bytes) -> None: + self.sent_frames.append(frame) + if len(self.sent_frames) >= self.disconnect_after: + raise WebSocketDisconnect(code=1000) + +class _CaptureThenDisconnectWebSocket: + def __init__(self) -> None: + self.sent_frames: list[bytes] = [] + + async def send_bytes(self, frame: bytes) -> None: + self.sent_frames.append(frame) + raise WebSocketDisconnect(code=1000) + + +async def _capture_one_agent_frame(frame: bytes, **ws_out_kwargs) -> _CaptureThenDisconnectWebSocket: + ws = _CaptureThenDisconnectWebSocket() + agent_q: asyncio.Queue[bytes] = asyncio.Queue() + await agent_q.put(frame) + + await asyncio.wait_for( + ws_out_loop( + ws, + agent_q, + frame_ms=20, + bytes_per_frame=len(frame), + **ws_out_kwargs, + ), + timeout=0.2, + ) + + return ws + + +def test_ws_out_loop_exits_when_close_was_already_sent() -> None: + async def _run() -> None: + ws = _ClosingWebSocket(RuntimeError('Cannot call "send" once a close message has been sent.')) + + await asyncio.wait_for( + ws_out_loop( + ws, + asyncio.Queue(), + frame_ms=20, + bytes_per_frame=4, + ), + timeout=0.2, + ) + + assert ws.sent_frames == 1 + + asyncio.run(_run()) + + +def test_ws_out_loop_applies_default_output_gain(monkeypatch) -> None: + monkeypatch.delenv("WS_OUTPUT_GAIN", raising=False) + + frame = struct.pack(" None: + monkeypatch.setenv("WS_OUTPUT_GAIN", "1.5") + + frame = struct.pack(" None: + monkeypatch.setenv("WS_OUTPUT_GAIN", "1.5") + + frame = struct.pack(" None: + async def _run() -> None: + ws = _ClosingWebSocket(WebSocketDisconnect(code=1000)) + + await asyncio.wait_for( + ws_out_loop( + ws, + asyncio.Queue(), + frame_ms=20, + bytes_per_frame=4, + ), + timeout=0.2, + ) + + assert ws.sent_frames == 1 + + asyncio.run(_run()) + + +def test_ws_out_loop_waits_for_speech_before_signaling_first_agent_audio() -> None: + async def _run() -> None: + silent_frame = struct.pack(" None: + logger = logging.getLogger("test.call_timeline") + logger.handlers = [] + logger.addHandler(logging.NullHandler()) + + with tempfile.TemporaryDirectory() as tmpdir: + old_dir = os.environ.get("CALL_TIMELINE_DIR") + old_console = os.environ.get("CALL_TIMELINE_CONSOLE") + old_enabled = os.environ.get("CALL_TIMELINE_ENABLED") + + os.environ["CALL_TIMELINE_DIR"] = tmpdir + os.environ["CALL_TIMELINE_CONSOLE"] = "0" + os.environ["CALL_TIMELINE_ENABLED"] = "1" + + try: + timeline = CallTimeline( + logger=logger, + component="bridge", + timeline_id="room-dev-123", + protocol="PRT-1", + room="room-dev-123", + session_id="PRT-1", + phone_number="31999999999", + origin_unix_ms=1000, + ) + timeline.emit("call_start", ok=True, nested={"a": 1}) + timeline.emit("ready_sent") + self.assertTrue(timeline.flush(timeout=5.0)) + + with timeline.path.open("r", encoding="utf-8") as handle: + lines = [json.loads(line) for line in handle if line.strip()] + finally: + if old_dir is None: + os.environ.pop("CALL_TIMELINE_DIR", None) + else: + os.environ["CALL_TIMELINE_DIR"] = old_dir + + if old_console is None: + os.environ.pop("CALL_TIMELINE_CONSOLE", None) + else: + os.environ["CALL_TIMELINE_CONSOLE"] = old_console + + if old_enabled is None: + os.environ.pop("CALL_TIMELINE_ENABLED", None) + else: + os.environ["CALL_TIMELINE_ENABLED"] = old_enabled + + self.assertEqual(2, len(lines)) + self.assertEqual("bridge", lines[0]["component"]) + self.assertEqual("call_start", lines[0]["event"]) + self.assertEqual("room-dev-123", lines[0]["timeline_id"]) + self.assertEqual("PRT-1", lines[0]["protocol"]) + self.assertEqual("31999999999", lines[0]["phone_number"]) + self.assertTrue(lines[0]["ok"]) + self.assertEqual({"a": 1}, lines[0]["nested"]) + self.assertGreaterEqual(lines[0]["t_rel_ms"], 0) + self.assertEqual("ready_sent", lines[1]["event"]) + + def test_background_writer_preserves_order_for_many_events(self) -> None: + logger = logging.getLogger("test.call_timeline.order") + logger.handlers = [logging.NullHandler()] + + with tempfile.TemporaryDirectory() as tmpdir: + old_dir = os.environ.get("CALL_TIMELINE_DIR") + old_enabled = os.environ.get("CALL_TIMELINE_ENABLED") + os.environ["CALL_TIMELINE_DIR"] = tmpdir + os.environ["CALL_TIMELINE_ENABLED"] = "1" + try: + timeline = CallTimeline( + logger=logger, + component="bridge", + timeline_id="room-order", + origin_unix_ms=1000, + ) + for index in range(200): + timeline.emit("tick", index=index) + self.assertTrue(timeline.flush(timeout=5.0)) + with timeline.path.open("r", encoding="utf-8") as handle: + lines = [json.loads(line) for line in handle if line.strip()] + finally: + if old_dir is None: + os.environ.pop("CALL_TIMELINE_DIR", None) + else: + os.environ["CALL_TIMELINE_DIR"] = old_dir + if old_enabled is None: + os.environ.pop("CALL_TIMELINE_ENABLED", None) + else: + os.environ["CALL_TIMELINE_ENABLED"] = old_enabled + + self.assertEqual(list(range(200)), [line["index"] for line in lines]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/utils/test_call_timeline_async_writer.py b/tests/utils/test_call_timeline_async_writer.py new file mode 100644 index 0000000..8cc1437 --- /dev/null +++ b/tests/utils/test_call_timeline_async_writer.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +import tempfile +import time +from pathlib import Path + +from app.utils.call_timeline import _TimelineWriter + + +def test_queue_full_is_counted_and_flush_honors_timeout() -> None: + writer = _TimelineWriter(max_queue=1, warning_interval_s=3_600) + writer.start = lambda: True # type: ignore[method-assign] + writer._last_drop_warning_at = time.monotonic() + + assert writer.submit(Path("unused"), "first") is True + assert writer.submit(Path("unused"), "second") is False + + started = time.monotonic() + assert writer.flush(timeout=0.01) is False + assert time.monotonic() - started < 0.5 + assert writer.dropped == 1 + + +def test_write_error_is_counted_without_killing_writer() -> None: + writer = _TimelineWriter(max_queue=10, warning_interval_s=3_600) + writer._last_error_warning_at = time.monotonic() + with tempfile.TemporaryDirectory() as tmpdir: + invalid_path = Path(tmpdir) + assert writer.submit(invalid_path, "line") is True + assert writer.flush(timeout=2.0) is True + + assert writer.write_errors == 1 + assert writer.shutdown(timeout=2.0) is True + + +def test_shutdown_drains_pending_lines_and_rejects_new_work() -> None: + writer = _TimelineWriter(max_queue=10) + writer._last_drop_warning_at = time.monotonic() + with tempfile.TemporaryDirectory() as tmpdir: + path = Path(tmpdir) / "timeline.jsonl" + assert writer.submit(path, "one") is True + assert writer.submit(path, "two") is True + assert writer.shutdown(timeout=2.0) is True + assert path.read_text(encoding="utf-8").splitlines() == ["one", "two"] + assert writer.submit(path, "three") is False + + assert writer.dropped == 1 diff --git a/tests/utils/test_full_call_recording.py b/tests/utils/test_full_call_recording.py new file mode 100644 index 0000000..d2121d8 --- /dev/null +++ b/tests/utils/test_full_call_recording.py @@ -0,0 +1,292 @@ +from __future__ import annotations + +import asyncio +import logging +import struct +import wave +from datetime import datetime +from pathlib import Path +from types import SimpleNamespace + +from app.utils import full_call_recording +from app.utils.full_call_recording import ( + EntireCallRecorder, + EntireCallUploadItem, + build_entire_call_object_name, + enqueue_entire_call_upload, +) +from app.utils.stt_audio_upload import OCIUploadConfig + + +def test_entire_call_object_name_uses_date_folder_and_session_filename() -> None: + object_name = build_entire_call_object_name( + session_id="sessao muito longa/" * 30, + now=datetime(2026, 7, 8, 10, 30, 0), + ) + + parts = object_name.split("/") + assert parts[0] == "2026-07-08" + assert parts[1] == "entire_call" + assert object_name.endswith(".wav") + assert len(parts) == 3 + assert len(parts[2]) <= 68 + assert " " not in object_name + + +def test_entire_call_recorder_writes_stereo_wav_without_upload(tmp_path: Path) -> None: + async def _run() -> Path: + recorder = EntireCallRecorder( + session_id="session-1", + sample_rate=1000, + channels=1, + sample_width=2, + frame_ms=20, + bytes_per_frame=4, + tmp_dir=tmp_path, + logger_override=logging.getLogger("test.full_call_recording.wav"), + object_name="2026-07-08/entire_call/session-1.wav", + ) + recorder.start() + recorder.record_client_frame(struct.pack(" None: + async def _run() -> Path: + recorder = EntireCallRecorder( + session_id="session-separated", + sample_rate=1000, + channels=1, + sample_width=2, + frame_ms=20, + bytes_per_frame=4, + tmp_dir=tmp_path, + logger_override=logging.getLogger("test.full_call_recording.separated"), + ) + recorder.start() + recorder.record_output_frame(struct.pack(" None: + async def _run() -> None: + recorder = EntireCallRecorder( + session_id="session-1", + sample_rate=1000, + channels=1, + sample_width=2, + frame_ms=20, + bytes_per_frame=4, + tmp_dir=tmp_path, + logger_override=logging.getLogger("test.full_call_recording.metadata"), + metadata={"session_id": "session-1", "room": "room-1"}, + object_name="2026-07-08/entire_call/session-1.wav", + ) + + recorder.start() + recorder.record_output_frame(b"\x00\x00\x00\x00") + result = await recorder.finalize(enqueue_upload=False) + + assert result is not None + + asyncio.run(_run()) + + +def test_entire_call_recorder_queue_full_drops_without_raising(tmp_path: Path) -> None: + recorder = EntireCallRecorder( + session_id="session-1", + sample_rate=1000, + channels=1, + sample_width=2, + frame_ms=20, + bytes_per_frame=4, + tmp_dir=tmp_path, + logger_override=logging.getLogger("test.full_call_recording.drop"), + queue_size=1, + ) + recorder._started = True + recorder._queue.put_nowait(object()) + + recorder.record_client_frame(b"\x00\x00\x00\x00") + + assert recorder.dropped_queue_frames == 1 + + +def test_upload_item_streams_file_to_oci(monkeypatch, tmp_path: Path) -> None: + calls: dict[str, object] = {} + wav_path = tmp_path / "call.wav" + wav_path.write_bytes(b"fake-wav") + config = OCIUploadConfig( + auth_mode="oke_workload_identity", + region="sa-saopaulo-1", + bucket="tia-audio", + namespace="namespace", + ) + + class FakeClient: + def put_object(self, **kwargs): + calls["namespace_name"] = kwargs["namespace_name"] + calls["bucket_name"] = kwargs["bucket_name"] + calls["object_name"] = kwargs["object_name"] + calls["body"] = kwargs["put_object_body"].read() + return SimpleNamespace(status=200) + + monkeypatch.setattr(full_call_recording, "_oci_upload_config_from_env", lambda: config) + monkeypatch.setattr(full_call_recording, "_oci_client", lambda _config: FakeClient()) + + item = EntireCallUploadItem( + object_name="2026-07-08/entire_call/session-1.wav", + path=wav_path, + session_id="session-1", + bytes=8, + duration_ms=20, + dropped_queue_frames=0, + dropped_client_buffer_frames=0, + metadata={}, + ) + + asyncio.run(full_call_recording._upload_item(item)) + + assert calls == { + "namespace_name": "namespace", + "bucket_name": "tia-audio", + "object_name": "2026-07-08/entire_call/session-1.wav", + "body": b"fake-wav", + } + + +def test_upload_item_retries_with_a_fresh_client_after_transient_failure( + monkeypatch, + tmp_path: Path, +) -> None: + wav_path = tmp_path / "call.wav" + wav_path.write_bytes(b"fake-wav") + config = OCIUploadConfig( + auth_mode="oke_workload_identity", + region="sa-saopaulo-1", + bucket="tia-audio", + namespace="namespace", + ) + attempts = 0 + resets = 0 + + class FakeClient: + def put_object(self, **_kwargs): + nonlocal attempts + attempts += 1 + if attempts == 1: + raise RuntimeError("transient connection failure") + return SimpleNamespace(status=200) + + def fake_reset() -> None: + nonlocal resets + resets += 1 + + monkeypatch.setenv("ENTIRE_CALL_RECORDING_UPLOAD_RETRY_BASE_DELAY_S", "0") + monkeypatch.setattr(full_call_recording, "_oci_upload_config_from_env", lambda: config) + monkeypatch.setattr(full_call_recording, "_oci_client", lambda _config: FakeClient()) + monkeypatch.setattr(full_call_recording, "_reset_oci_client", fake_reset) + + item = EntireCallUploadItem( + object_name="2026-07-08/entire_call/session-retry.wav", + path=wav_path, + session_id="session-retry", + bytes=8, + duration_ms=20, + dropped_queue_frames=0, + dropped_client_buffer_frames=0, + metadata={}, + ) + + asyncio.run(full_call_recording._upload_item(item)) + + assert attempts == 2 + assert resets == 1 + + +def test_upload_worker_keeps_file_after_all_attempts_fail(monkeypatch, tmp_path: Path) -> None: + async def _run() -> None: + wav_path = tmp_path / "failed.wav" + wav_path.write_bytes(b"recoverable-wav") + item = EntireCallUploadItem( + object_name="2026-07-08/entire_call/session-failed.wav", + path=wav_path, + session_id="session-failed", + bytes=15, + duration_ms=20, + dropped_queue_frames=0, + dropped_client_buffer_frames=0, + metadata={}, + ) + + async def fail_upload(_item) -> None: + raise RuntimeError("OCI unavailable") + + monkeypatch.setattr(full_call_recording, "_upload_item", fail_upload) + worker = full_call_recording._EntireCallUploadWorker(asyncio.get_running_loop()) + assert worker.enqueue(item) + await asyncio.wait_for(worker._queue.join(), timeout=1) + + assert wav_path.read_bytes() == b"recoverable-wav" + worker._task.cancel() + + asyncio.run(_run()) + + +def test_enqueue_entire_call_upload_is_noop_and_keeps_file_without_oci_config( + monkeypatch, + tmp_path: Path, +) -> None: + wav_path = tmp_path / "call.wav" + wav_path.write_bytes(b"fake-wav") + monkeypatch.setattr(full_call_recording, "_oci_upload_config_from_env", lambda: None) + + assert ( + enqueue_entire_call_upload( + path=wav_path, + object_name="2026-07-08/entire_call/session-1.wav", + session_id="session-1", + bytes=8, + duration_ms=20, + logger_override=logging.getLogger("test.full_call_recording.noop"), + ) + is None + ) + assert wav_path.read_bytes() == b"fake-wav" diff --git a/tests/utils/test_logging_async_file.py b/tests/utils/test_logging_async_file.py new file mode 100644 index 0000000..cc041d5 --- /dev/null +++ b/tests/utils/test_logging_async_file.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +import logging +import os +import io +import tempfile +import time +from pathlib import Path + +import app.utils.logging as app_logging +from app.utils.logging import _LogWriter, get_call_logger + + +class _FailingHandler(logging.Handler): + def emit(self, record: logging.LogRecord) -> None: + raise OSError("disk unavailable") + + +def test_session_console_formatter_appends_context_without_duplicating() -> None: + stream = io.StringIO() + handler = logging.StreamHandler(stream) + handler.setFormatter(app_logging._SessionConsoleFormatter("%(message)s")) + logger = logging.getLogger(f"test.session.console.{time.time_ns()}") + logger.handlers = [handler] + logger.propagate = False + logger.setLevel(logging.INFO) + + try: + app_logging.set_log_session_id("session-123") + logger.info("FLOW | step=test") + logger.info("FLOW | step=test | session_id=session-123") + finally: + app_logging.set_log_session_id("") + + assert stream.getvalue().splitlines() == [ + "FLOW | step=test | session_id=session-123", + "FLOW | step=test | session_id=session-123", + ] + + +def test_session_console_formatter_keeps_structured_json_valid() -> None: + formatter = app_logging._SessionConsoleFormatter("%(message)s") + try: + app_logging.set_log_session_id("session-456") + record = logging.LogRecord( + "test", + logging.INFO, + "", + 0, + '{"tipo_evento":"envio msg","session_id":"session-456"}', + (), + None, + ) + + assert formatter.format(record) == ( + '{"tipo_evento":"envio msg","session_id":"session-456"}' + ) + finally: + app_logging.set_log_session_id("") + + +def test_call_file_handler_is_async_and_writes_to_disk() -> None: + with tempfile.TemporaryDirectory() as tmp: + saved = {key: os.environ.get(key) for key in ("LOG_DIR", "LOG_TO_FILE")} + os.environ["LOG_DIR"] = tmp + os.environ["LOG_TO_FILE"] = "1" + try: + logger = get_call_logger( + phone_number="5500000000123", + session_id=f"async{time.time_ns()}", + ) + async_handlers = [ + handler + for handler in logger.handlers + if isinstance(handler, app_logging._AsyncFileHandler) + ] + assert len(async_handlers) == 1 + + logger.info("LINHA_DE_TESTE_ASYNC_123") + assert app_logging._LOG_WRITER.flush(timeout=5.0) + + files = list(Path(tmp).glob("*.txt")) + assert len(files) == 1 + assert "LINHA_DE_TESTE_ASYNC_123" in files[0].read_text(encoding="utf-8") + finally: + for key, value in saved.items(): + if value is None: + os.environ.pop(key, None) + else: + os.environ[key] = value + + +def test_queue_full_is_counted_and_flush_honors_timeout() -> None: + writer = _LogWriter(max_queue=1, warning_interval_s=3_600) + writer.start = lambda: True # type: ignore[method-assign] + writer._last_drop_warning_at = time.monotonic() + handler = logging.NullHandler() + record = logging.LogRecord("test", logging.INFO, "", 0, "message", (), None) + + assert writer.submit(handler, record) is True + assert writer.submit(handler, record) is False + + started = time.monotonic() + assert writer.flush(timeout=0.01) is False + assert time.monotonic() - started < 0.5 + assert writer.dropped == 1 + + +def test_handler_error_is_counted_and_shutdown_stays_healthy() -> None: + writer = _LogWriter(max_queue=10, warning_interval_s=3_600) + writer._last_error_warning_at = time.monotonic() + record = logging.LogRecord("test", logging.INFO, "", 0, "message", (), None) + + assert writer.submit(_FailingHandler(), record) is True + assert writer.flush(timeout=2.0) is True + assert writer.write_errors == 1 + assert writer.shutdown(timeout=2.0) is True diff --git a/tests/utils/test_logging_pubsub.py b/tests/utils/test_logging_pubsub.py new file mode 100644 index 0000000..8e320b0 --- /dev/null +++ b/tests/utils/test_logging_pubsub.py @@ -0,0 +1,460 @@ +from __future__ import annotations + +import json +import logging +import uuid + +from app.utils import logging as logging_utils +from app.utils import structured_otlp +from app.utils import structured_pubsub + + +class _FakeFuture: + def add_done_callback(self, callback): + callback(self) + + def result(self): + return "message-id" + + +class _FakePublisher: + def __init__(self) -> None: + self.calls = [] + + def publish(self, *args, **kwargs): + self.calls.append((args, kwargs)) + return _FakeFuture() + + def topic_path(self, project_id: str, topic: str) -> str: + return f"projects/{project_id}/topics/{topic}" + + +def test_log_structured_event_publishes_to_pubsub_when_configured(monkeypatch) -> None: + published_events = [] + + monkeypatch.setenv("STRUCTURED_EVENT_LOG_ENABLED", "1") + monkeypatch.setattr(logging_utils, "publish_structured_event", published_events.append) + + logger = logging.getLogger("test.logging_pubsub") + logger.handlers = [logging.NullHandler()] + logger.setLevel(logging.INFO) + + event = logging_utils.log_structured_event( + logger, + logging_utils.StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ), + tipo_evento=logging_utils.EVENT_RECEBIMENTO_MSG, + message_id="message-1", + inicio_ns=1_700_000_000_000_000_000, + ) + + assert event is not None + assert published_events == [event] + + +def test_log_structured_event_publishes_to_pubsub_and_otlp_when_trace_id_is_valid(monkeypatch) -> None: + published_events = [] + published_spans = [] + + monkeypatch.setenv("STRUCTURED_EVENT_LOG_ENABLED", "1") + monkeypatch.setattr(logging_utils, "publish_structured_event", published_events.append) + monkeypatch.setattr(logging_utils, "publish_structured_span", published_spans.append) + + logger = logging.getLogger("test.logging_pubsub_otlp") + logger.handlers = [logging.NullHandler()] + logger.setLevel(logging.INFO) + + event = logging_utils.log_structured_event( + logger, + logging_utils.StructuredLogContext( + callid="call-1", + session_id="550e8400-e29b-41d4-a716-446655440101", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ), + tipo_evento=logging_utils.EVENT_RECEBIMENTO_MSG, + message_id="message-1", + inicio_ns=1_700_000_000_000_000_000, + ) + + assert event is not None + assert published_events == [event] + assert published_spans == [event] + + +def test_publish_structured_event_sends_json_to_configured_topic(monkeypatch) -> None: + fake_publisher = _FakePublisher() + + monkeypatch.setenv("GCP_PROJECT_ID", "project-1") + monkeypatch.setenv("AGENT_PUBSUB_TOPIC", "agent-logs") + monkeypatch.setattr(structured_pubsub, "_PUBLISHER", fake_publisher) + monkeypatch.setattr(structured_pubsub, "_TOPIC_PATH", "projects/project-1/topics/agent-logs") + + event = { + "tipo_evento": "recebimento msg", + "dat_hora_inicio": "2026-05-11T22:07:29.965Z", + "dat_hora_fim": "2026-05-11T22:07:30.965Z", + "callid": "call-1", + "session_id": "session-1", + "nome_agente": "conta", + "cod_ani": "1234", + "message_id": "message-1", + "original_message_id": "original-message-1", + "http_cod_status": 200, + "http_cod_desc": "OK", + } + structured_pubsub.publish_structured_event(event) + + assert len(fake_publisher.calls) == 1 + + args, kwargs = fake_publisher.calls[0] + assert args[0] == "projects/project-1/topics/agent-logs" + payload = json.loads(args[1].decode("utf-8")) + assert payload == { + "tipo_evento": "recebimento msg", + "dat_hora_inicio": "11/05/2026 19:07:29,965000000", + "dat_hora_termino": "11/05/2026 19:07:30,965000000", + "callid": "call-1", + "sessionId": "session-1", + "nome_agente": "conta", + "cod_ani": "1234", + "message_id": "message-1", + "http_cod_status": 200, + "http_cod_desc": "OK", + } + assert "session_id" not in payload + assert "dat_hora_fim" not in payload + assert "original_message_id" not in payload + assert kwargs == { + "tipo_evento": "recebimento msg", + "dat_hora_inicio": "11/05/2026 19:07:29,965000000", + "dat_hora_termino": "11/05/2026 19:07:30,965000000", + "sessionId": "session-1", + "callid": "call-1", + "nome_agente": "conta", + "cod_ani": "1234", + "http_cod_status": "200", + "http_cod_desc": "OK", + } + + +def test_log_structured_event_skips_pubsub_without_required_env(monkeypatch) -> None: + fake_publisher = _FakePublisher() + + monkeypatch.delenv("GCP_PROJECT_ID", raising=False) + monkeypatch.delenv("AGENT_PUBSUB_TOPIC", raising=False) + monkeypatch.setattr(structured_pubsub, "_PUBLISHER", fake_publisher) + monkeypatch.setattr(structured_pubsub, "_TOPIC_PATH", "projects/project-1/topics/agent-logs") + + structured_pubsub.publish_structured_event({"tipo_evento": "envio msg"}) + + assert fake_publisher.calls == [] + + +def test_build_structured_event_includes_session_id_and_lowercase_nome_agente() -> None: + event = logging_utils.build_structured_event( + logging_utils.StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ), + tipo_evento=logging_utils.EVENT_RECEBIMENTO_MSG, + message_id="message-1", + inicio_ns=1_700_000_000_000_000_000, + ) + + assert event["session_id"] == "session-1" + assert event["nome_agente"] == "conta" + assert event["dat_hora_fim"] == "" + assert event["latencia_total_STT_TTS"] == "" + assert event["latencia_TFFB_STT_TTS"] == "" + assert event["duracao_audio"] == "" + assert event["tts_max_gap_ms"] == "" + assert event["tts_max_underrun_0ms"] == "" + assert event["tts_underflow_count"] == "" + assert event["tts_avg_underflow_ms"] == "" + assert "tts_clear_rtt_ms" not in event + assert "tts_connection_queue_wait_ms" not in event + assert "tts_clear_stale_messages" not in event + assert "tts_clear_stale_audio_bytes" not in event + assert event["interrupcao"] == "" + assert event["erro_msg"] == "" + assert event["erro_detalhe"] == "" + assert event["http_cod_status"] == "" + assert event["http_cod_desc"] == "" + assert event["finalizacao"] == "" + assert all(value is not None for value in event.values()) + assert "Nome_agente" not in event + + +def test_build_structured_event_includes_audio_duration_and_interruption_flag() -> None: + event = logging_utils.build_structured_event( + logging_utils.StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ), + tipo_evento=logging_utils.EVENT_ENVIO_MSG, + message_id="message-1", + inicio_ns=1_700_000_000_000_000_000, + duracao_audio_ms=2400, + tts_max_gap_ms=840, + tts_max_underrun_0ms=125, + interrupcao=True, + ) + + assert event["duracao_audio"] == 2400 + assert event["tts_max_gap_ms"] == 840 + assert event["tts_max_underrun_0ms"] == 125 + assert event["interrupcao"] == 1 + + +def test_build_structured_event_uses_uuid_when_message_id_is_missing(monkeypatch) -> None: + generated = uuid.UUID("12345678-1234-4234-9234-123456789abc") + monkeypatch.setattr(logging_utils.uuid, "uuid4", lambda: generated) + + event = logging_utils.build_structured_event( + logging_utils.StructuredLogContext( + callid="call-1", + session_id="session-1", + num_telefone="5511999999999", + cod_ani="1234", + nome_agente="conta", + ), + tipo_evento=logging_utils.EVENT_ENVIO_MSG, + inicio_ns=1_700_000_000_000_000_000, + ) + + assert event["message_id"] == "12345678-1234-4234-9234-123456789abc" + + +def test_error_message_from_resource_uses_event_specific_contract() -> None: + assert ( + logging_utils.error_message_from_resource( + resource="agent_runtime", + tipo_evento=logging_utils.EVENT_RECEBIMENTO_MSG, + ) + == "Falha TIA" + ) + assert ( + logging_utils.error_message_from_resource( + resource="stt", + tipo_evento=logging_utils.EVENT_RECEBIMENTO_MSG, + ) + == "Falha STT" + ) + assert ( + logging_utils.error_message_from_resource( + resource="bridge", + tipo_evento=logging_utils.EVENT_RECEBIMENTO_MSG, + ) + == "Falha comunicacao" + ) + assert ( + logging_utils.error_message_from_resource( + status="stop_silencio_longo", + tipo_evento=logging_utils.EVENT_RECEBIMENTO_MSG, + ) + == "Silencio Longo" + ) + assert ( + logging_utils.error_message_from_resource( + resource="agent_backend", + tipo_evento=logging_utils.EVENT_ENVIO_MSG, + ) + == "Falha comunicacao" + ) + assert ( + logging_utils.error_message_from_resource( + resource="tts", + tipo_evento=logging_utils.EVENT_ENVIO_MSG, + ) + == "Falha TTS" + ) + assert ( + logging_utils.error_message_from_resource( + status="transferred", + tipo_evento=logging_utils.EVENT_ENVIO_MSG, + ) + == "Transferido" + ) + + +def test_format_event_timestamp_uses_iso_utc_milliseconds() -> None: + assert logging_utils.format_event_timestamp(1_700_000_001_234_567_890) == ( + "2023-11-14T22:13:21.234Z" + ) + + +def test_structured_context_from_start_data_keeps_session_id_empty_when_not_provided() -> None: + context = logging_utils.structured_context_from_start_data( + { + "callIdGed": "ged-1", + "gsm": "5511999999999", + "ani": "1234", + "agent": "conta", + } + ) + + assert context.callid == "ged-1" + assert context.session_id == "" + + +def test_topic_path_accepts_short_topic_or_full_resource_name() -> None: + fake_publisher = _FakePublisher() + + assert structured_pubsub._topic_path(fake_publisher, "project-1", "agent-logs") == ( + "projects/project-1/topics/agent-logs" + ) + assert structured_pubsub._topic_path( + fake_publisher, + "project-1", + "projects/another-project/topics/agent-logs", + ) == "projects/another-project/topics/agent-logs" + + +def test_trace_id_from_session_id_sanitizes_uuid_and_rejects_invalid_values() -> None: + assert structured_otlp.trace_id_from_session_id( + "550e8400-e29b-41d4-a716-446655440101" + ) == "550e8400e29b41d4a716446655440101" + assert structured_otlp.trace_id_from_session_id("sess-001") == "" + assert structured_otlp.trace_id_from_session_id("00000000-0000-0000-0000-000000000000") == "" + + +def _install_memory_tracer(monkeypatch): + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + + endpoint = "http://otel.example/v1/traces" + exporter = InMemorySpanExporter() + provider = TracerProvider(id_generator=structured_otlp._ID_GENERATOR) + provider.add_span_processor(SimpleSpanProcessor(exporter)) + + monkeypatch.setenv("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT", endpoint) + monkeypatch.setattr(structured_otlp, "_TRACER_PROVIDER", provider) + monkeypatch.setattr(structured_otlp, "_TRACER", provider.get_tracer("test.structured_otlp")) + monkeypatch.setattr(structured_otlp, "_TRACER_ENDPOINT", endpoint) + return exporter + + +def test_publish_structured_span_exports_otlp_span_with_timestamps_and_attributes(monkeypatch) -> None: + exporter = _install_memory_tracer(monkeypatch) + start_ns = 1_700_000_000_000_000_000 + end_ns = 1_700_000_001_234_000_000 + + event = { + "tipo_evento": "envio msg", + "message_id": "message-1", + "dat_hora_inicio": logging_utils.format_event_timestamp(start_ns), + "dat_hora_fim": logging_utils.format_event_timestamp(end_ns), + "callid": "call-1", + "session_id": "550e8400-e29b-41d4-a716-446655440101", + "num_telefone": "5511999999999", + "cod_ani": "1234", + "latencia_total_STT_TTS": 1234, + "latencia_TFFB_STT_TTS": 123, + "erro_msg": "Falha TTS", + "erro_detalhe": "", + "http_cod_status": "", + "http_cod_desc": "", + "finalizacao": "", + "nome_agente": "conta", + "original_message_id": "original-message-1", + } + + structured_otlp.publish_structured_span(event) + + spans = exporter.get_finished_spans() + assert len(spans) == 1 + span = spans[0] + assert span.name == "structured_log.envio msg" + assert span.context.trace_id == int("550e8400e29b41d4a716446655440101", 16) + assert span.context.span_id != 0 + assert span.parent is None + assert span.start_time == start_ns + assert span.end_time == end_ns + assert span.attributes["tipo_evento"] == "envio msg" + assert span.attributes["message_id"] == "message-1" + assert span.attributes["latencia_total_STT_TTS"] == 1234 + assert span.attributes["erro_msg"] == "Falha TTS" + assert span.attributes["erro_detalhe"] == "" + assert span.attributes["http_cod_status"] == "" + assert span.attributes["http_cod_desc"] == "" + assert span.attributes["finalizacao"] == "" + assert "original_message_id" not in span.attributes + assert span.status.status_code.name == "ERROR" + + +def test_publish_structured_span_keeps_same_trace_id_and_lets_sdk_generate_span_id(monkeypatch) -> None: + exporter = _install_memory_tracer(monkeypatch) + trace_id = int("550e8400e29b41d4a716446655440101", 16) + + base_event = { + "dat_hora_inicio": logging_utils.format_event_timestamp(1_700_000_000_000_000_000), + "session_id": "550e8400-e29b-41d4-a716-446655440101", + } + structured_otlp.publish_structured_span({**base_event, "tipo_evento": "recebimento msg"}) + structured_otlp.publish_structured_span({**base_event, "tipo_evento": "envio msg"}) + + spans = exporter.get_finished_spans() + assert len(spans) == 2 + assert {span.context.trace_id for span in spans} == {trace_id} + assert spans[0].context.span_id != spans[1].context.span_id + assert spans[0].parent is None + assert spans[1].parent is None + + +def test_publish_structured_span_omits_empty_tffb_attribute(monkeypatch) -> None: + exporter = _install_memory_tracer(monkeypatch) + + structured_otlp.publish_structured_span( + { + "tipo_evento": "envio msg", + "dat_hora_inicio": logging_utils.format_event_timestamp(1_700_000_000_000_000_000), + "session_id": "550e8400-e29b-41d4-a716-446655440101", + "latencia_TFFB_STT_TTS": "", + } + ) + + spans = exporter.get_finished_spans() + assert len(spans) == 1 + assert "latencia_TFFB_STT_TTS" not in spans[0].attributes + + +def test_publish_structured_span_skips_invalid_session_id(monkeypatch) -> None: + exporter = _install_memory_tracer(monkeypatch) + + structured_otlp.publish_structured_span( + { + "tipo_evento": "recebimento msg", + "dat_hora_inicio": logging_utils.format_event_timestamp(1_700_000_000_000_000_000), + "session_id": "sess-001", + } + ) + + assert exporter.get_finished_spans() == () + + +def test_publish_structured_span_skips_without_otlp_endpoint(monkeypatch) -> None: + monkeypatch.delenv("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT", raising=False) + monkeypatch.setattr(structured_otlp, "_TRACER", None) + monkeypatch.setattr(structured_otlp, "_TRACER_ENDPOINT", "") + + structured_otlp.publish_structured_span( + { + "tipo_evento": "recebimento msg", + "dat_hora_inicio": logging_utils.format_event_timestamp(1_700_000_000_000_000_000), + "session_id": "550e8400-e29b-41d4-a716-446655440101", + } + ) diff --git a/tests/utils/test_structured_otlp.py b/tests/utils/test_structured_otlp.py new file mode 100644 index 0000000..fe863fc --- /dev/null +++ b/tests/utils/test_structured_otlp.py @@ -0,0 +1,85 @@ +from __future__ import annotations + +import os + +from opentelemetry.sdk.trace.export import SpanExportResult + +import app.utils.structured_otlp as structured_otlp + + +class _RecordingExporter: + def __init__(self, *args: object, **kwargs: object) -> None: + self.exported: list[int] = [] + self.shutdown_calls = 0 + + def export(self, spans: object) -> SpanExportResult: + self.exported.append(len(list(spans))) + return SpanExportResult.SUCCESS + + def shutdown(self) -> None: + self.shutdown_calls += 1 + + def force_flush(self, timeout_millis: int = 30_000) -> bool: + return True + + +def _reset_provider_without_shutdown() -> None: + with structured_otlp._TRACER_LOCK: + structured_otlp._TRACER = None + structured_otlp._TRACER_PROVIDER = None + structured_otlp._TRACER_ENDPOINT = "" + + +def test_span_export_is_batched_not_synchronous(monkeypatch) -> None: + old_endpoint = os.environ.get("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT") + os.environ["OTEL_EXPORTER_OTLP_TRACES_ENDPOINT"] = "http://localhost:4318/v1/traces" + exporter = _RecordingExporter() + monkeypatch.setattr( + structured_otlp, + "OTLPSpanExporter", + lambda *args, **kwargs: exporter, + ) + _reset_provider_without_shutdown() + try: + structured_otlp.publish_structured_span( + { + "session_id": "0123456789abcdef0123456789abcdef", + "tipo_evento": "audio_in", + "dat_hora_inicio": "2026-07-24T12:00:00.000000Z", + "dat_hora_fim": "2026-07-24T12:00:00.010000Z", + } + ) + + assert exporter.exported == [] + assert structured_otlp.force_flush_structured_otlp(timeout_millis=5_000) + assert sum(exporter.exported) >= 1 + finally: + structured_otlp.shutdown_structured_otlp() + if old_endpoint is None: + os.environ.pop("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT", None) + else: + os.environ["OTEL_EXPORTER_OTLP_TRACES_ENDPOINT"] = old_endpoint + + +def test_shutdown_flushes_provider_once(monkeypatch) -> None: + old_endpoint = os.environ.get("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT") + os.environ["OTEL_EXPORTER_OTLP_TRACES_ENDPOINT"] = "http://localhost:4318/v1/traces" + exporter = _RecordingExporter() + monkeypatch.setattr( + structured_otlp, + "OTLPSpanExporter", + lambda *args, **kwargs: exporter, + ) + _reset_provider_without_shutdown() + try: + assert structured_otlp._tracer() is not None + structured_otlp.shutdown_structured_otlp() + structured_otlp.shutdown_structured_otlp() + assert exporter.shutdown_calls == 1 + assert structured_otlp._TRACER_PROVIDER is None + finally: + structured_otlp.shutdown_structured_otlp() + if old_endpoint is None: + os.environ.pop("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT", None) + else: + os.environ["OTEL_EXPORTER_OTLP_TRACES_ENDPOINT"] = old_endpoint diff --git a/tests/utils/test_stt_audio_upload.py b/tests/utils/test_stt_audio_upload.py new file mode 100644 index 0000000..30ce0b4 --- /dev/null +++ b/tests/utils/test_stt_audio_upload.py @@ -0,0 +1,221 @@ +from __future__ import annotations + +import sys +import types +from datetime import datetime + +from app.utils import stt_audio_upload +from app.utils.stt_audio_upload import ( + build_stt_vad_object_name, + enqueue_stt_vad_audio_upload, +) + + +_BUCKET_ENV_NAMES = ( + "OCI_AUTH_MODE", + "BUCKET_USER_ID", + "BUCKET_PRIVATE_KEY", + "BUCKET_FINGERPRINT", + "BUCKET_TENANCY_ID", + "BUCKET_REGION", + "BUCKET_NAME", + "BUCKET_NAMESPACE", +) + + +def _clear_bucket_env(monkeypatch) -> None: + for name in _BUCKET_ENV_NAMES: + monkeypatch.delenv(name, raising=False) + + +def _install_fake_oci_modules(monkeypatch) -> dict[str, object]: + calls: dict[str, object] = {} + + class FakeObjectStorageClient: + def __init__(self, config, **kwargs) -> None: + calls["client_config"] = config + calls["client_kwargs"] = kwargs + + def fake_oke_signer() -> object: + signer = object() + calls["signer"] = signer + return signer + + fake_oci = types.ModuleType("oci") + fake_auth = types.ModuleType("oci.auth") + fake_signers = types.ModuleType("oci.auth.signers") + fake_object_storage = types.ModuleType("oci.object_storage") + + fake_object_storage.ObjectStorageClient = FakeObjectStorageClient + fake_signers.get_oke_workload_identity_resource_principal_signer = fake_oke_signer + fake_auth.signers = fake_signers + fake_oci.auth = fake_auth + fake_oci.object_storage = fake_object_storage + + monkeypatch.setitem(sys.modules, "oci", fake_oci) + monkeypatch.setitem(sys.modules, "oci.auth", fake_auth) + monkeypatch.setitem(sys.modules, "oci.auth.signers", fake_signers) + monkeypatch.setitem(sys.modules, "oci.object_storage", fake_object_storage) + return calls + + +def _reset_oci_client_cache(monkeypatch) -> None: + monkeypatch.setattr(stt_audio_upload, "_OCI_CLIENT", None) + monkeypatch.setattr(stt_audio_upload, "_OCI_CLIENT_CONFIG", None) + + +def test_stt_vad_object_name_uses_requested_layout_and_stays_short() -> None: + object_name = build_stt_vad_object_name( + session_id="sessao muito longa/" * 30, + message_id="message-id-muito-longo/" * 30, + req_id="req-id-muito-longo/" * 30, + now=datetime(2026, 6, 24, 13, 52, 21), + ) + + parts = object_name.split("/") + assert parts[0] == "2026-06-24" + assert object_name.endswith(".wav") + assert len(parts) == 3 + assert len(object_name.encode("utf-8")) <= 512 + assert all(len(part) <= 68 for part in parts[1:]) + assert " " not in object_name + + +def test_stt_vad_object_name_uses_message_id_as_filename() -> None: + object_name = build_stt_vad_object_name( + session_id="417a5a59-961a-4ba6-9c68-71e92a3da46d", + message_id="0bc6875e-ebd1-4dca-a43d-2f7de404895b", + req_id="c182775c8c114e3f833bfb7675e0be70", + now=datetime(2026, 6, 29, 10, 30, 0), + ) + + assert ( + object_name + == "2026-06-29/417a5a59-961a-4ba6-9c68-71e92a3da46d/" + "0bc6875e-ebd1-4dca-a43d-2f7de404895b.wav" + ) + assert "c182775c8c114e3f833bfb7675e0be70" not in object_name + + +def test_oci_upload_config_reads_bucket_env_and_private_key(monkeypatch) -> None: + _clear_bucket_env(monkeypatch) + monkeypatch.setenv("OCI_AUTH_MODE", "local") + monkeypatch.setenv("BUCKET_USER_ID", "user-ocid") + monkeypatch.setenv("BUCKET_PRIVATE_KEY", "-----BEGIN KEY-----\\nabc\\n-----END KEY-----") + monkeypatch.setenv("BUCKET_FINGERPRINT", "fingerprint") + monkeypatch.setenv("BUCKET_TENANCY_ID", "tenancy-ocid") + monkeypatch.setenv("BUCKET_REGION", "sa-saopaulo-1") + monkeypatch.setenv("BUCKET_NAME", "tia-audio") + monkeypatch.setenv("BUCKET_NAMESPACE", "namespace") + + config = stt_audio_upload._oci_upload_config_from_env() + + assert config is not None + assert config.auth_mode == "local" + assert config.bucket == "tia-audio" + assert config.namespace == "namespace" + assert config.key_content == "-----BEGIN KEY-----\nabc\n-----END KEY-----" + + +def test_oci_upload_config_uses_oke_when_auth_mode_is_absent(monkeypatch) -> None: + _clear_bucket_env(monkeypatch) + monkeypatch.setenv("BUCKET_REGION", "sa-saopaulo-1") + monkeypatch.setenv("BUCKET_NAME", "tia-audio") + monkeypatch.setenv("BUCKET_NAMESPACE", "namespace") + + config = stt_audio_upload._oci_upload_config_from_env() + + assert config is not None + assert config.auth_mode == "oke_workload_identity" + assert config.bucket == "tia-audio" + assert config.namespace == "namespace" + assert config.user == "" + assert config.key_content == "" + + +def test_oci_upload_config_uses_oke_when_auth_mode_is_not_local(monkeypatch) -> None: + _clear_bucket_env(monkeypatch) + monkeypatch.setenv("OCI_AUTH_MODE", "prd") + monkeypatch.setenv("BUCKET_REGION", "sa-saopaulo-1") + monkeypatch.setenv("BUCKET_NAME", "tia-audio") + monkeypatch.setenv("BUCKET_NAMESPACE", "namespace") + + config = stt_audio_upload._oci_upload_config_from_env() + + assert config is not None + assert config.auth_mode == "oke_workload_identity" + + +def test_oci_upload_config_is_missing_when_local_required_env_is_absent(monkeypatch) -> None: + _clear_bucket_env(monkeypatch) + monkeypatch.setenv("OCI_AUTH_MODE", "local") + monkeypatch.setenv("BUCKET_REGION", "sa-saopaulo-1") + monkeypatch.setenv("BUCKET_NAME", "tia-audio") + monkeypatch.setenv("BUCKET_NAMESPACE", "namespace") + + assert stt_audio_upload._oci_upload_config_from_env() is None + + +def test_oci_upload_config_is_missing_when_oke_required_env_is_absent(monkeypatch) -> None: + _clear_bucket_env(monkeypatch) + monkeypatch.setenv("BUCKET_REGION", "sa-saopaulo-1") + monkeypatch.setenv("BUCKET_NAME", "tia-audio") + + assert stt_audio_upload._oci_upload_config_from_env() is None + + +def test_oci_client_uses_api_key_config_for_local_mode(monkeypatch) -> None: + _reset_oci_client_cache(monkeypatch) + calls = _install_fake_oci_modules(monkeypatch) + config = stt_audio_upload.OCIUploadConfig( + auth_mode="local", + region="sa-saopaulo-1", + bucket="tia-audio", + namespace="namespace", + user="user-ocid", + key_content="private-key", + fingerprint="fingerprint", + tenancy="tenancy-ocid", + ) + + stt_audio_upload._oci_client(config) + + assert calls["client_config"] == { + "user": "user-ocid", + "key_content": "private-key", + "fingerprint": "fingerprint", + "tenancy": "tenancy-ocid", + "region": "sa-saopaulo-1", + } + assert calls["client_kwargs"] == {} + assert "signer" not in calls + + +def test_oci_client_uses_oke_signer_for_workload_identity(monkeypatch) -> None: + _reset_oci_client_cache(monkeypatch) + calls = _install_fake_oci_modules(monkeypatch) + config = stt_audio_upload.OCIUploadConfig( + auth_mode="oke_workload_identity", + region="sa-saopaulo-1", + bucket="tia-audio", + namespace="namespace", + ) + + stt_audio_upload._oci_client(config) + + assert calls["client_config"] == {"region": "sa-saopaulo-1"} + assert calls["client_kwargs"] == {"signer": calls["signer"]} + + +def test_enqueue_stt_vad_audio_upload_is_noop_without_oci_config(monkeypatch) -> None: + monkeypatch.setattr(stt_audio_upload, "_oci_upload_config_from_env", lambda: None) + + assert ( + enqueue_stt_vad_audio_upload( + wav_bytes=b"fake-wav", + req_id="req-1", + message_id="message-1", + structured_log_context={"session_id": "session-1"}, + ) + is None + ) diff --git a/tests/utils/test_turn_ids.py b/tests/utils/test_turn_ids.py new file mode 100644 index 0000000..56f6b0e --- /dev/null +++ b/tests/utils/test_turn_ids.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +import uuid +from unittest import mock + +from app.utils.turn_ids import ( + clear_started_turn_message_id, + next_turn_message_id, + peek_started_turn_message_id, + register_started_turn_message_id, + reset_turn_message_sequence, +) + + +def test_next_turn_message_id_uses_uuid4() -> None: + generated = uuid.UUID("12345678-1234-4234-9234-123456789abc") + + with mock.patch("app.utils.turn_ids.uuid.uuid4", return_value=generated): + message_id = next_turn_message_id( + { + "callIdGed": "GED-555", + "protocol_id": "PRT-555", + "session_id": "session-555", + } + ) + + assert message_id == "12345678-1234-4234-9234-123456789abc" + + +def test_started_turn_message_id_can_be_registered_peeked_and_cleared() -> None: + source = {"session_id": "session-started-turn"} + reset_turn_message_sequence(source, clear_pending=True) + + assert peek_started_turn_message_id(source) == "" + + register_started_turn_message_id(source, "message-1") + assert peek_started_turn_message_id(source) == "message-1" + + register_started_turn_message_id(source, "message-2") + assert peek_started_turn_message_id(source) == "message-2" + + clear_started_turn_message_id(source, "message-1") + assert peek_started_turn_message_id(source) == "message-2" + + clear_started_turn_message_id(source, "message-2") + assert peek_started_turn_message_id(source) == "" + + +def test_reset_turn_message_sequence_clears_started_turn_message_id() -> None: + source = {"session_id": "session-reset-started-turn"} + reset_turn_message_sequence(source, clear_pending=True) + register_started_turn_message_id(source, "message-to-reset") + + reset_turn_message_sequence(source, clear_pending=True) + + assert peek_started_turn_message_id(source) == "" diff --git a/tests/ws_gateway/__init__.py b/tests/ws_gateway/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/ws_gateway/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/ws_gateway/test_audio_flow_logging.py b/tests/ws_gateway/test_audio_flow_logging.py new file mode 100644 index 0000000..70ead6f --- /dev/null +++ b/tests/ws_gateway/test_audio_flow_logging.py @@ -0,0 +1,9 @@ +from app.ws_gateway import main as main_module + + +def test_audio_flow_burst_gap_controls_repeated_logs(monkeypatch) -> None: + monkeypatch.setattr(main_module, "FLOW_AUDIO_BURST_GAP_S", 0.8) + + assert main_module._is_new_audio_burst(10.0, 0.0) is True + assert main_module._is_new_audio_burst(10.5, 10.0) is False + assert main_module._is_new_audio_burst(10.8, 10.0) is True diff --git a/tests/ws_gateway/test_audio_input_backlog.py b/tests/ws_gateway/test_audio_input_backlog.py new file mode 100644 index 0000000..3f32460 --- /dev/null +++ b/tests/ws_gateway/test_audio_input_backlog.py @@ -0,0 +1,305 @@ +import asyncio +import struct + +from app.ws_gateway import main as main_module + + +def _frame(value: int) -> bytes: + return bytes([value % 256]) * main_module.BYTES_PER_FRAME + + +def _silent_frame() -> bytes: + return b"\x00" * main_module.BYTES_PER_FRAME + + +def _voiced_frame(amp: int = 8000) -> bytes: + samples = main_module.BYTES_PER_FRAME // 2 + return struct.pack("<%dh" % samples, *([amp] * samples)) + + +def _config( + *, + enabled: bool = True, + energy_shed_enabled: bool = False, + energy_shed_max_excess_ms: int = 1000, + silence_dbfs: float = -50.0, +) -> main_module.AudioInputBacklogConfig: + return main_module.AudioInputBacklogConfig( + shed_enabled=enabled, + shed_threshold_ms=100, + shed_keep_ms=40, + latency_metrics_enabled=True, + latency_alert_ms=1000, + latency_log_interval_s=15.0, + livekit_source_queue_size_ms=500, + livekit_source_clear_on_shed=True, + energy_shed_enabled=energy_shed_enabled, + energy_shed_max_excess_ms=energy_shed_max_excess_ms, + silence_dbfs=silence_dbfs, + ) + + +def test_audio_input_latency_snapshot_separates_raw_and_effective_delay() -> None: + tracker = main_module.AudioInputLatencyTracker(_config()) + tracker.mark_enabled(100.0) + + fields = tracker.snapshot(frame_count=1064, now=110.872, queue_size_frames=0) + + assert fields["received_audio_ms"] == 21280 + assert fields["elapsed_since_enable_ms"] == 10872 + assert fields["raw_excess_ms"] == 10408 + assert fields["raw_lag_ms"] == 10408 + assert fields["excess_ms"] == 10408 + assert fields["lag_ms"] == 10408 + + tracker.total_dropped_frames = 520 + fields = tracker.snapshot(frame_count=1064, now=110.872, queue_size_frames=0) + + assert fields["total_dropped_ms"] == 10400 + assert fields["saved_latency_ms"] == 10400 + assert fields["raw_excess_ms"] == 10408 + assert fields["raw_lag_ms"] == 10408 + assert fields["excess_ms"] == 8 + assert fields["lag_ms"] == 8 + + +def test_shed_audio_queue_backlog_keeps_recent_frames_and_logs_savings(monkeypatch) -> None: + events = [] + debug_events = [] + + def fake_log_flow_event(logger, step, **payload): + events.append((step, payload)) + + monkeypatch.setattr(main_module, "log_flow_event", fake_log_flow_event) + audio_q: asyncio.Queue[bytes] = asyncio.Queue() + for value in range(8): + audio_q.put_nowait(_frame(value)) + tracker = main_module.AudioInputLatencyTracker( + _config(), + debug_event_publisher=lambda event, payload: debug_events.append( + (event, payload) + ), + ) + tracker.mark_enabled(0.0) + + dropped = main_module._shed_audio_queue_backlog( + audio_q, + config=_config(), + tracker=tracker, + source="test", + frame_count=8, + now=1.0, + ) + + assert dropped.dropped == 6 + assert audio_q.qsize() == 2 + assert [audio_q.get_nowait(), audio_q.get_nowait()] == [_frame(6), _frame(7)] + assert events[-1][0] == "audio_in_latency_shed" + assert events[-1][1]["dropped_ms"] == 120 + assert events[-1][1]["saved_ms"] == 120 + assert events[-1][1]["saved_latency_ms"] == 120 + assert events[-1][1]["queue_before_ms"] == 160 + assert events[-1][1]["queue_after_ms"] == 40 + assert events[-1][1]["total_dropped_ms"] == 120 + assert debug_events[-1][0] == "bridge.audio_in.shed" + assert debug_events[-1][1]["dropped_ms"] == 120 + + +def test_shed_audio_queue_backlog_is_disabled_by_flag(monkeypatch) -> None: + events = [] + monkeypatch.setattr( + main_module, + "log_flow_event", + lambda logger, step, **payload: events.append((step, payload)), + ) + audio_q: asyncio.Queue[bytes] = asyncio.Queue() + for value in range(8): + audio_q.put_nowait(_frame(value)) + + dropped = main_module._shed_audio_queue_backlog( + audio_q, + config=_config(enabled=False), + tracker=main_module.AudioInputLatencyTracker(_config(enabled=False)), + source="test", + frame_count=8, + now=1.0, + ) + + assert dropped.dropped == 0 + assert audio_q.qsize() == 8 + assert events == [] + + +def test_audio_input_backlog_config_defaults_to_enabled(monkeypatch) -> None: + for name in ( + "AUDIO_IN_BACKLOG_SHED_ENABLED", + "AUDIO_IN_BACKLOG_SHED_THRESHOLD_MS", + "AUDIO_IN_BACKLOG_SHED_KEEP_MS", + "AUDIO_IN_LATENCY_METRICS_ENABLED", + "AUDIO_IN_LATENCY_ALERT_MS", + "AUDIO_IN_LATENCY_LOG_INTERVAL_S", + "LIVEKIT_AUDIO_SOURCE_QUEUE_SIZE_MS", + "LIVEKIT_AUDIO_SOURCE_CLEAR_ON_SHED", + ): + monkeypatch.delenv(name, raising=False) + + config = main_module.audio_input_backlog_config_from_env() + + assert config.shed_enabled is True + assert config.shed_threshold_ms == 500 + assert config.shed_keep_ms == 300 + assert config.latency_metrics_enabled is True + assert config.livekit_source_queue_size_ms == 500 + assert config.livekit_source_clear_on_shed is True + assert config.config_source == "env" + assert config.call_config_overrides == () + + +def test_audio_input_backlog_config_preserves_legacy_source_queue_when_disabled(monkeypatch) -> None: + monkeypatch.setenv("AUDIO_IN_BACKLOG_SHED_ENABLED", "0") + monkeypatch.delenv("LIVEKIT_AUDIO_SOURCE_QUEUE_SIZE_MS", raising=False) + monkeypatch.delenv("LIVEKIT_AUDIO_SOURCE_CLEAR_ON_SHED", raising=False) + + config = main_module.audio_input_backlog_config_from_env() + + assert config.shed_enabled is False + assert config.livekit_source_queue_size_ms == 5000 + assert config.livekit_source_clear_on_shed is False + assert config.config_source == "env" + + +def test_audio_input_backlog_config_prefers_call_config_and_falls_back_to_env(monkeypatch) -> None: + monkeypatch.setenv("AUDIO_IN_BACKLOG_SHED_ENABLED", "0") + monkeypatch.setenv("AUDIO_IN_BACKLOG_SHED_THRESHOLD_MS", "900") + monkeypatch.setenv("AUDIO_IN_BACKLOG_SHED_KEEP_MS", "600") + monkeypatch.setenv("AUDIO_IN_LATENCY_METRICS_ENABLED", "1") + monkeypatch.setenv("AUDIO_IN_LATENCY_ALERT_MS", "1200") + monkeypatch.setenv("AUDIO_IN_LATENCY_LOG_INTERVAL_S", "20") + monkeypatch.setenv("LIVEKIT_AUDIO_SOURCE_QUEUE_SIZE_MS", "2000") + monkeypatch.setenv("LIVEKIT_AUDIO_SOURCE_CLEAR_ON_SHED", "0") + + config = main_module.audio_input_backlog_config_from_env( + { + "ws": { + "audioInputBacklogShedEnabled": True, + "audioInputBacklogShedThresholdMs": 500, + "audioInputBacklogShedKeepMs": 300, + "audioInputLatencyAlertMs": 800, + "livekitAudioSourceQueueSizeMs": 400, + "livekitAudioSourceClearOnShed": True, + "audioInputBacklogEnergyShedEnabled": True, + "audioInputBacklogEnergyShedMaxExcessMs": 600, + "audioInputBacklogSilenceDbfs": -60, + } + } + ) + + assert config.shed_enabled is True + assert config.shed_threshold_ms == 500 + assert config.shed_keep_ms == 300 + assert config.latency_metrics_enabled is True + assert config.latency_alert_ms == 800 + assert config.latency_log_interval_s == 20.0 + assert config.livekit_source_queue_size_ms == 400 + assert config.livekit_source_clear_on_shed is True + assert config.config_source == "call_config" + assert config.call_config_overrides == ( + "audio_in_backlog_shed_enabled", + "audio_in_backlog_shed_threshold_ms", + "audio_in_backlog_shed_keep_ms", + "audio_in_latency_alert_ms", + "livekit_audio_source_queue_size_ms", + "livekit_audio_source_clear_on_shed", + "audio_in_backlog_energy_shed_enabled", + "audio_in_backlog_energy_shed_max_excess_ms", + "audio_in_backlog_silence_dbfs", + ) + + +def test_shed_energy_guided_drops_only_silence_under_small_excess(monkeypatch) -> None: + events = [] + monkeypatch.setattr( + main_module, + "log_flow_event", + lambda logger, step, **payload: events.append((step, payload)), + ) + config = _config(energy_shed_enabled=True, energy_shed_max_excess_ms=1000) + # [V, S, V, S, V, S, V, S] -> 4 de voz, 4 de silencio + frames = [] + for _ in range(4): + frames.append(_voiced_frame()) + frames.append(_silent_frame()) + audio_q: asyncio.Queue[bytes] = asyncio.Queue() + for frame in frames: + audio_q.put_nowait(frame) + tracker = main_module.AudioInputLatencyTracker(config) + tracker.mark_enabled(0.0) + + dropped = main_module._shed_audio_queue_backlog( + audio_q, + config=config, + tracker=tracker, + source="test", + frame_count=8, + now=1.0, + ) + + # excesso pequeno (120 ms < 1000): descarta so os 4 silencios, preserva a voz. + assert dropped.dropped == 4 + assert audio_q.qsize() == 4 + remaining = [audio_q.get_nowait() for _ in range(4)] + assert all(f == _voiced_frame() for f in remaining) + assert events[-1][0] == "audio_in_latency_shed" + assert events[-1][1]["mode"] == "energy" + assert events[-1][1]["dropped_silent_frames"] == 4 + assert events[-1][1]["dropped_voiced_frames"] == 0 + assert events[-1][1]["content_policy"] == "silence_only" + assert events[-1][1]["sync_action"] == "compress_silence" + + +def test_shed_falls_back_to_blind_under_large_excess(monkeypatch) -> None: + events = [] + monkeypatch.setattr( + main_module, + "log_flow_event", + lambda logger, step, **payload: events.append((step, payload)), + ) + # max_excess minusculo: qualquer excesso cai no regime de rajada (cego) + config = _config(energy_shed_enabled=True, energy_shed_max_excess_ms=40) + audio_q: asyncio.Queue[bytes] = asyncio.Queue() + for value in range(8): + audio_q.put_nowait(_frame(value)) + tracker = main_module.AudioInputLatencyTracker(config) + tracker.mark_enabled(0.0) + + dropped = main_module._shed_audio_queue_backlog( + audio_q, + config=config, + tracker=tracker, + source="test", + frame_count=8, + now=1.0, + ) + + assert dropped.dropped == 6 + assert audio_q.qsize() == 2 + assert [audio_q.get_nowait(), audio_q.get_nowait()] == [_frame(6), _frame(7)] + assert events[-1][1]["mode"] == "blind" + assert events[-1][1]["content_policy"] == "oldest_frames" + assert events[-1][1]["sync_action"] == "drop_audio_to_catch_up" + + +def test_audio_input_backlog_config_energy_defaults(monkeypatch) -> None: + for name in ( + "AUDIO_IN_BACKLOG_ENERGY_SHED_ENABLED", + "AUDIO_IN_BACKLOG_ENERGY_SHED_MAX_EXCESS_MS", + "AUDIO_IN_BACKLOG_SILENCE_DBFS", + ): + monkeypatch.delenv(name, raising=False) + + config = main_module.audio_input_backlog_config_from_env() + + assert config.energy_shed_enabled is True + assert config.energy_shed_max_excess_ms == 1000 + assert config.silence_dbfs == -60.0 + assert config.silence_rms_threshold == 32 diff --git a/tests/ws_gateway/test_fake_remote_agent.py b/tests/ws_gateway/test_fake_remote_agent.py new file mode 100644 index 0000000..7d7d5fa --- /dev/null +++ b/tests/ws_gateway/test_fake_remote_agent.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +import unittest + +from app.ws_gateway.fake_remote_agent import build_fake_remote_agent_response + + +CONTA_OPENING = ( + "Olá! Eu sou a Especialista em Contas e vou ajudar você a entender a sua " + "fatura. Posso explicar valores, detalhar serviços e itens eventuais, " + "identificar cobranças que você não reconhece e, se for o caso, realizar " + "ajustes necessários ou solicitações relacionadas à sua conta. Então vamos " + "lá, me conte o que você gostaria de entender ou resolver na sua conta." +) + + +class FakeRemoteAgentTests(unittest.TestCase): + def test_conta_opening_uses_the_configured_greeting(self) -> None: + response = build_fake_remote_agent_response( + { + "action": "chat", + "payload": {"agent": "conta", "stage": "PRESENTATION"}, + } + ) + + self.assertEqual("ARGUMENTATION", response["stage"]) + self.assertEqual(CONTA_OPENING, response["result"]["content"]) + + def test_conta_contract_returns_result_content(self) -> None: + response = build_fake_remote_agent_response( + { + "action": "chat", + "payload": { + "agent": "conta", + "stage": "PRESENTATION", + "message": "quero detalhes da fatura", + "msisdn": "5511999990000", + }, + } + ) + + self.assertEqual("result", response["type"]) + self.assertEqual("ARGUMENTATION", response["stage"]) + self.assertEqual("final", response["result"]["type"]) + self.assertIn("fatura", response["result"]["content"].lower()) + self.assertIn("quero detalhes da fatura", response["result"]["content"].lower()) + + def test_conta_done_when_user_asks_to_finish(self) -> None: + response = build_fake_remote_agent_response( + { + "action": "chat", + "payload": { + "agent": "conta", + "stage": "FORMALIZATION", + "message": "obrigado, pode encerrar", + }, + } + ) + + self.assertEqual("DONE", response["stage"]) + self.assertEqual("result", response["type"]) + + def test_generic_contract_returns_text(self) -> None: + response = build_fake_remote_agent_response( + { + "agent": "oferta", + "stage": "PRESENTATION", + "text": "quero saber da oferta", + } + ) + + self.assertEqual("final", response["type"]) + self.assertEqual("ARGUMENTATION", response["stage"]) + self.assertIn("oferta", response["text"].lower()) + self.assertIn("quero saber da oferta", response["text"].lower()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/ws_gateway/test_readiness.py b/tests/ws_gateway/test_readiness.py new file mode 100644 index 0000000..77ec17d --- /dev/null +++ b/tests/ws_gateway/test_readiness.py @@ -0,0 +1,1056 @@ +from __future__ import annotations + +import asyncio +import json +from types import SimpleNamespace +from unittest import mock + +from fastapi.testclient import TestClient + +from app.ws_gateway import main as main_module +from app.ws_gateway.readiness import ( + BridgeReadinessManager, + CapacityReservation, + ReadinessReport, + ResourceCheck, + ServicesHealthReport, + build_bridge_failed_stop_message, + build_capacity_stop_message, + build_completed_stop_message, +) + + +def _build_start_payload() -> dict: + return { + "type": "start", + "data": { + "agent": "conta", + "ani": "5511999999999", + "gsm": "5511999999999", + "session_id": "550e8400-e29b-41d4-a716-446655440200", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-001", + "agentData": { + "idFatura": "fat-123", + }, + }, + } + + +class _FakeLease: + def __init__(self) -> None: + self.release_calls = 0 + + async def release(self) -> None: + self.release_calls += 1 + + +class _FakeReadinessManager: + def __init__( + self, + *, + reservation: CapacityReservation | None = None, + report: ReadinessReport | None = None, + ) -> None: + self.lease = reservation.lease if reservation is not None else _FakeLease() + self.reservation = reservation or CapacityReservation( + allowed=True, + lease=self.lease, + active_connections=1, + max_connections=5, + ) + self.report = report or ReadinessReport( + checks={ + "agent_runtime": ResourceCheck(status="ok"), + "agent_backend": ResourceCheck(status="ok"), + "stt": ResourceCheck(status="ok"), + "tts": ResourceCheck(status="ok"), + }, + failed_resources=[], + active_connections=1, + max_connections=5, + cached=False, + ) + self.services_report = ServicesHealthReport( + checks={ + "agent_runtime": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + "agent_backend": { + "contas": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + "oferta": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + }, + "stt": { + "sofya": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + }, + "tts": { + "xAI": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + }, + }, + healthy=True, + ) + self.evaluate_calls: list[tuple[str, dict]] = [] + self.default_calls = 0 + self.services_calls = 0 + + async def reserve_capacity(self) -> CapacityReservation: + return self.reservation + + async def evaluate_resources(self, *, agent_name: str, call_config: dict | None) -> ReadinessReport: + self.evaluate_calls.append((agent_name, dict(call_config or {}))) + return self.report + + async def evaluate_default_resources(self) -> ReadinessReport: + self.default_calls += 1 + return self.report + + async def evaluate_services(self) -> ServicesHealthReport: + self.services_calls += 1 + return self.services_report + + +class _FakeTimeline: + def __init__(self) -> None: + self.events: list[tuple[str, dict]] = [] + self.path = "/tmp/timeline.jsonl" + + def emit(self, event: str, **fields) -> None: + self.events.append((event, fields)) + + +def test_stop_message_helpers_follow_contract() -> None: + assert build_capacity_stop_message() == { + "type": "stop", + "data": { + "status": "stop_capacity_tia", + "reason": "capacity_exceeded", + "resource": "capacity", + "failed_resources": ["capacity"], + "phase": "pre_ready", + }, + } + assert build_completed_stop_message("finished") == { + "type": "stop", + "data": { + "status": "stop_resolvido_e_finalizado", + "reason": "finished", + "phase": "in_session", + }, + } + assert build_bridge_failed_stop_message() == { + "type": "stop", + "data": { + "status": "stop_bridge_failed", + "reason": "bridge_failed", + "resource": "bridge", + "phase": "in_session", + }, + } + + +def test_build_completed_stop_message_maps_reason_to_configured_status(monkeypatch) -> None: + monkeypatch.setenv("FINAL_STOP_STATUS_RESOLVED", "stop_custom_resolvido") + monkeypatch.setenv("FINAL_STOP_STATUS_UNRESOLVED", "stop_custom_nao_resolvido") + monkeypatch.setenv("FINAL_STOP_STATUS_OTHER_SUBJECT", "stop_custom_outro_assunto") + monkeypatch.setenv("FINAL_STOP_STATUS_LONG_SILENCE", "stop_custom_silencio") + monkeypatch.setenv("FINAL_STOP_REASON_RESOLVED", "resolvido") + monkeypatch.setenv("FINAL_STOP_REASON_UNRESOLVED", "nao_resolvido") + monkeypatch.setenv("FINAL_STOP_REASON_OTHER_SUBJECT", "outro_assunto") + monkeypatch.setenv("FINAL_STOP_REASON_LONG_SILENCE", "silencio_longo") + + assert build_completed_stop_message("nao_resolvido") == { + "type": "stop", + "data": { + "status": "stop_custom_nao_resolvido", + "reason": "nao_resolvido", + "phase": "in_session", + }, + } + assert build_completed_stop_message("outro_assunto") == { + "type": "stop", + "data": { + "status": "stop_custom_outro_assunto", + "reason": "outro_assunto", + "phase": "in_session", + }, + } + assert build_completed_stop_message("silencio_longo") == { + "type": "stop", + "data": { + "status": "stop_custom_silencio", + "reason": "silencio_longo", + "phase": "in_session", + }, + } + assert build_completed_stop_message("resolvido") == { + "type": "stop", + "data": { + "status": "stop_custom_resolvido", + "reason": "resolvido", + "phase": "in_session", + }, + } + + +def test_capacity_reservation_blocks_and_releases(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("TIA_WS_MAX_CONNECTIONS", "1") + + first = await manager.reserve_capacity() + second = await manager.reserve_capacity() + assert first.allowed is True + assert second.allowed is False + + await first.lease.release() + + third = await manager.reserve_capacity() + assert third.allowed is True + await third.lease.release() + + asyncio.run(_run()) + + +def test_probe_agent_backend_prefers_explicit_health_url(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("REMOTE_AGENT_HEALTH_URL_CONTA", "https://agent.example/health") + monkeypatch.setenv("AGENT_BACKEND", "remote_ws") + + async def _fake_probe(check_name: str, url: str) -> ResourceCheck: + assert check_name == "agent_backend" + assert url == "https://agent.example/health" + return ResourceCheck(status="ok", url=url) + + with mock.patch.object(manager, "_probe_http_health", side_effect=_fake_probe): + check = await manager.probe_agent_backend(agent_name="conta", call_config={}) + + assert check.status == "ok" + assert check.url == "https://agent.example/health" + + asyncio.run(_run()) + + +def test_probe_agent_backend_derives_health_url_from_websocket(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL_CONTA", raising=False) + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL", raising=False) + monkeypatch.setenv("REMOTE_AGENT_WS_URL", "wss://agent.internal/agent/ws") + monkeypatch.setenv("AGENT_BACKEND", "remote_ws") + + async def _fake_probe(check_name: str, url: str) -> ResourceCheck: + assert check_name == "agent_backend" + assert url == "https://agent.internal/health" + return ResourceCheck(status="ok", url=url) + + with mock.patch.object(manager, "_probe_http_health", side_effect=_fake_probe): + check = await manager.probe_agent_backend(agent_name="conta", call_config={}) + + assert check.status == "ok" + assert check.url == "https://agent.internal/health" + + asyncio.run(_run()) + + +def test_probe_agent_backend_derives_health_url_from_sse(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL_CONTA", raising=False) + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL_CONTAS", raising=False) + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL", raising=False) + monkeypatch.setenv("REMOTE_AGENT_SSE_URL_CONTA", "http://agent.internal/agent/sse") + monkeypatch.setenv("AGENT_BACKEND", "remote_sse") + + async def _fake_probe(check_name: str, url: str) -> ResourceCheck: + assert check_name == "agent_backend" + assert url == "http://agent.internal/health" + return ResourceCheck(status="ok", url=url) + + with mock.patch.object(manager, "_probe_http_health", side_effect=_fake_probe): + check = await manager.probe_agent_backend(agent_name="conta", call_config={}) + + assert check.status == "ok" + assert check.url == "http://agent.internal/health" + + asyncio.run(_run()) + + +def test_probe_agent_backend_uses_oferta_health_url(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("REMOTE_AGENT_HEALTH_URL_OFERTA", "https://oferta.example/health") + monkeypatch.setenv("AGENT_BACKEND", "remote_sse") + + async def _fake_probe(check_name: str, url: str) -> ResourceCheck: + assert check_name == "agent_backend" + assert url == "https://oferta.example/health" + return ResourceCheck(status="ok", url=url) + + with mock.patch.object(manager, "_probe_http_health", side_effect=_fake_probe): + check = await manager.probe_agent_backend(agent_name="oferta", call_config={}) + + assert check.status == "ok" + assert check.url == "https://oferta.example/health" + + asyncio.run(_run()) + + +def test_probe_agent_backend_skips_oferta_health_probe_without_explicit_route(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL_OFERTA", raising=False) + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL_OFERTAS", raising=False) + monkeypatch.setenv("REMOTE_AGENT_HEALTH_URL", "http://conta.internal/health") + monkeypatch.setenv("REMOTE_AGENT_SSE_URL_OFERTA", "http://oferta.internal/agent/execute") + monkeypatch.setenv("AGENT_BACKEND", "remote_sse") + + with mock.patch.object(manager, "_probe_http_health") as mocked_probe: + check = await manager.probe_agent_backend(agent_name="ofertas", call_config={}) + + mocked_probe.assert_not_called() + assert check.status == "skipped" + assert check.url == "" + assert "oferta backend health probe disabled" in check.message + + asyncio.run(_run()) + + +def test_probe_stt_derives_health_url_from_stt_url(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.delenv("STT_HEALTH_URL", raising=False) + monkeypatch.setenv("STT_PROVIDER", "internal_http") + monkeypatch.setenv("STT_URL", "http://stt.internal/api/transcriber") + + async def _fake_probe(check_name: str, url: str) -> ResourceCheck: + assert check_name == "stt" + assert url == "http://stt.internal/health" + return ResourceCheck(status="ok", url=url) + + with mock.patch.object(manager, "_probe_http_health", side_effect=_fake_probe): + check = await manager.probe_stt({}) + + assert check.status == "ok" + assert check.url == "http://stt.internal/health" + + asyncio.run(_run()) + + +def test_probe_stt_can_be_temporarily_skipped_by_env(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("TIA_SKIP_STT_READINESS", "1") + monkeypatch.setenv("STT_PROVIDER", "internal_http") + monkeypatch.setenv("STT_URL", "http://stt.internal/api/transcriber") + + with mock.patch.object(manager, "_probe_http_health") as mocked_probe: + check = await manager.probe_stt({}) + + assert check.status == "skipped" + assert check.message == "STT readiness disabled by env" + mocked_probe.assert_not_called() + + asyncio.run(_run()) + + +def test_probe_http_health_preserves_returned_status_and_description() -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + + class _FakeResponse: + status_code = 503 + reason_phrase = "Service Unavailable" + text = '{"message":"db down"}' + + def json(self) -> dict: + return {"message": "db down"} + + class _FakeAsyncClient: + def __init__(self, *args, **kwargs) -> None: + pass + + async def __aenter__(self) -> "_FakeAsyncClient": + return self + + async def __aexit__(self, *args) -> None: + return None + + async def get(self, url: str, timeout: float): + assert url == "https://service.example/health" + assert timeout > 0 + return _FakeResponse() + + with mock.patch("app.ws_gateway.readiness.httpx.AsyncClient", _FakeAsyncClient): + check = await manager._probe_http_health("agent_backend", "https://service.example/health") + + assert check.status == "fail" + assert check.message == "agent_backend returned status 503: db down" + assert check.http_status_code == 503 + assert check.http_status_desc == "db down" + + asyncio.run(_run()) + + +def test_evaluate_resources_uses_cache(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("TIA_RESOURCE_HEALTH_TTL_S", "60") + counters = { + "agent_runtime": 0, + "agent_backend": 0, + "stt": 0, + "tts": 0, + } + + async def _ok(name: str) -> ResourceCheck: + counters[name] += 1 + return ResourceCheck(status="ok") + + async def _probe_agent_runtime() -> ResourceCheck: + return await _ok("agent_runtime") + + async def _probe_agent_backend(**_: dict) -> ResourceCheck: + return await _ok("agent_backend") + + async def _probe_stt(*_: dict) -> ResourceCheck: + return await _ok("stt") + + async def _probe_tts(*_: dict) -> ResourceCheck: + return await _ok("tts") + + with ( + mock.patch.object(manager, "probe_agent_runtime", side_effect=_probe_agent_runtime), + mock.patch.object(manager, "probe_agent_backend", side_effect=_probe_agent_backend), + mock.patch.object(manager, "probe_stt", side_effect=_probe_stt), + mock.patch.object(manager, "probe_tts", side_effect=_probe_tts), + ): + first = await manager.evaluate_resources(agent_name="conta", call_config={}) + second = await manager.evaluate_resources(agent_name="conta", call_config={}) + + assert first.cached is False + assert second.cached is True + assert counters == { + "agent_runtime": 1, + "agent_backend": 1, + "stt": 1, + "tts": 1, + } + + asyncio.run(_run()) + + +def test_evaluate_resources_runs_resource_probes_in_parallel(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("TIA_RESOURCE_HEALTH_TTL_S", "0") + started: set[str] = set() + all_started = asyncio.Event() + + async def _probe(name: str) -> ResourceCheck: + started.add(name) + if len(started) == 4: + all_started.set() + await asyncio.wait_for(all_started.wait(), timeout=0.2) + return ResourceCheck(status="ok") + + async def _probe_agent_runtime() -> ResourceCheck: + return await _probe("agent_runtime") + + async def _probe_agent_backend(**_: dict) -> ResourceCheck: + return await _probe("agent_backend") + + async def _probe_stt(*_: dict) -> ResourceCheck: + return await _probe("stt") + + async def _probe_tts(*_: dict) -> ResourceCheck: + return await _probe("tts") + + with ( + mock.patch.object(manager, "probe_agent_runtime", side_effect=_probe_agent_runtime), + mock.patch.object(manager, "probe_agent_backend", side_effect=_probe_agent_backend), + mock.patch.object(manager, "probe_stt", side_effect=_probe_stt), + mock.patch.object(manager, "probe_tts", side_effect=_probe_tts), + ): + report = await manager.evaluate_resources(agent_name="conta", call_config={}) + + assert report.is_healthy() is True + assert started == {"agent_runtime", "agent_backend", "stt", "tts"} + + asyncio.run(_run()) + + +def test_probe_tts_fake_returns_ok() -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + check = await manager.probe_tts({"tts": {"provider": "fake"}}) + assert check.status == "ok" + + asyncio.run(_run()) + + +def test_probe_tts_elevenlabs_can_be_mocked(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("ELEVENLABS_API_KEY", "key-123") + monkeypatch.setenv("ELEVENLABS_VOICE_ID", "voice-123") + monkeypatch.setenv("ELEVENLABS_MODEL_ID", "model-123") + + fake_stream = mock.Mock() + fake_stream.collect = mock.AsyncMock(return_value=mock.Mock(data=b"\x00\x01")) + fake_tts = mock.Mock() + fake_tts.synthesize.return_value = fake_stream + fake_tts.aclose = mock.AsyncMock() + + with mock.patch("livekit.plugins.elevenlabs.TTS", return_value=fake_tts) as mocked_tts: + check = await manager.probe_tts({"tts": {"provider": "elevenlabs"}}) + + assert check.status == "ok" + assert mocked_tts.call_args.kwargs["http_session"] is not None + fake_tts.aclose.assert_awaited_once() + + asyncio.run(_run()) + + +def test_probe_tts_xai_can_be_mocked(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("XAI_API_KEY", "key-123") + + fake_stream = mock.Mock() + fake_stream.collect = mock.AsyncMock(return_value=mock.Mock(data=b"\x00\x01")) + fake_tts = mock.Mock() + fake_tts.synthesize.return_value = fake_stream + fake_tts.aclose = mock.AsyncMock() + + with mock.patch("app.ws_gateway.readiness.OraclexAITTS", return_value=fake_tts) as mocked_tts: + check = await manager.probe_tts( + {"tts": {"provider": "xai", "voiceId": "voice-123", "language": "pt-BR"}} + ) + + assert check.status == "ok" + assert mocked_tts.call_args.kwargs["api_key"] == "key-123" + assert mocked_tts.call_args.kwargs["voice"] == "voice-123" + assert mocked_tts.call_args.kwargs["language"] == "pt-BR" + assert mocked_tts.call_args.kwargs["http_session"] is not None + fake_tts.aclose.assert_awaited_once() + + asyncio.run(_run()) + + +def test_probe_tts_xai_can_only_connect(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("XAI_API_KEY", "key-123") + monkeypatch.setenv("XAI_TTS_READINESS_MODE", "connect") + + fake_tts = mock.Mock() + fake_tts.connect = mock.AsyncMock() + fake_tts.aclose = mock.AsyncMock() + + with mock.patch("app.ws_gateway.readiness.OraclexAITTS", return_value=fake_tts): + check = await manager.probe_tts({"tts": {"provider": "xai"}}) + + assert check.status == "ok" + assert check.message == "xAI TTS websocket connected" + fake_tts.connect.assert_awaited_once() + fake_tts.synthesize.assert_not_called() + fake_tts.aclose.assert_awaited_once() + + asyncio.run(_run()) + + +def test_probe_http_health_timeout_has_diagnostic_message(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("TIA_RESOURCE_HEALTH_TIMEOUT_S", "0.1") + + class _FailingAsyncClient: + def __init__(self, *args, **kwargs) -> None: + pass + + async def __aenter__(self) -> "_FailingAsyncClient": + return self + + async def __aexit__(self, *args) -> None: + return None + + async def get(self, url: str, timeout: float): + raise TimeoutError() + + with mock.patch("app.ws_gateway.readiness.httpx.AsyncClient", _FailingAsyncClient): + check = await manager._probe_http_health("stt", "https://service.example/health") + + assert check.status == "fail" + assert check.message == "stt readiness timed out after 100ms (TimeoutError)" + + asyncio.run(_run()) + + +def test_probe_tts_xai_timeout_has_diagnostic_message(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("XAI_API_KEY", "key-123") + monkeypatch.setenv("TIA_RESOURCE_HEALTH_TIMEOUT_S", "0.1") + + fake_tts = mock.Mock() + fake_tts.connect = mock.AsyncMock(side_effect=TimeoutError()) + fake_tts.aclose = mock.AsyncMock() + + with mock.patch("app.ws_gateway.readiness.OraclexAITTS", return_value=fake_tts): + monkeypatch.setenv("XAI_TTS_READINESS_MODE", "connect") + check = await manager.probe_tts({"tts": {"provider": "xai"}}) + + assert check.status == "fail" + assert check.message == "xai TTS readiness timed out after 100ms (TimeoutError)" + fake_tts.aclose.assert_awaited_once() + + asyncio.run(_run()) + + +def test_probe_tts_xai_iam_does_not_require_legacy_api_key(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("XAI_TTS_AUTH_METHOD", "INSTANCE_PRINCIPAL") + monkeypatch.setenv("OCI_COMPARTMENT_ID", "ocid1.compartment.oc1..example") + monkeypatch.setenv("XAI_TTS_READINESS_MODE", "connect") + monkeypatch.delenv("XAI_API_KEY", raising=False) + + fake_tts = mock.Mock() + fake_tts.connect = mock.AsyncMock() + fake_tts.aclose = mock.AsyncMock() + + with mock.patch("app.ws_gateway.readiness.OraclexAITTS", return_value=fake_tts) as mocked_tts: + check = await manager.probe_tts({"tts": {"provider": "xai"}}) + + assert check.status == "ok" + assert mocked_tts.call_args.kwargs["auth_method"] == "INSTANCE_PRINCIPAL" + assert "api_key" not in mocked_tts.call_args.kwargs + fake_tts.connect.assert_awaited_once() + fake_tts.aclose.assert_awaited_once() + + asyncio.run(_run()) + + +def test_probe_tts_xai_requires_api_key(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.delenv("XAI_API_KEY", raising=False) + + check = await manager.probe_tts({"tts": {"provider": "xai"}}) + + assert check.status == "fail" + assert check.message == "missing xAI TTS config: XAI_API_KEY" + + asyncio.run(_run()) + + +def test_probe_tts_azure_rest_can_be_mocked(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("TTS_PROVIDER", "azure") + monkeypatch.setenv("AZURE_TTS_IMPLEMENTATION", "rest") + monkeypatch.setenv("AZURE_SPEECH_KEY", "key-123") + monkeypatch.setenv("AZURE_SPEECH_REGION", "brazilsouth") + monkeypatch.setenv("AZURE_SPEECH_VOICE", "pt-BR-FranciscaNeural") + + with mock.patch("app.ws_gateway.readiness.AzureRESTTTS.synthesize_pcm", return_value=b"\x00\x01"): + check = await manager.probe_tts({"tts": {"provider": "azure"}}) + + assert check.status == "ok" + + asyncio.run(_run()) + + +def test_probe_tts_azure_plugin_can_be_mocked(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("TTS_PROVIDER", "azure") + monkeypatch.setenv("AZURE_TTS_IMPLEMENTATION", "plugin") + monkeypatch.setenv("AZURE_SPEECH_KEY", "key-123") + monkeypatch.setenv("AZURE_SPEECH_REGION", "brazilsouth") + monkeypatch.setenv("AZURE_SPEECH_VOICE", "pt-BR-FranciscaNeural") + + fake_stream = mock.Mock() + fake_stream.collect = mock.AsyncMock(return_value=mock.Mock(data=b"\x00\x01")) + fake_tts = mock.Mock() + fake_tts.synthesize.return_value = fake_stream + fake_tts.aclose = mock.AsyncMock() + + with mock.patch("livekit.plugins.azure.TTS", return_value=fake_tts) as mocked_tts: + check = await manager.probe_tts({"tts": {"provider": "azure"}}) + + assert check.status == "ok" + assert mocked_tts.call_args.kwargs["http_session"] is not None + fake_tts.aclose.assert_awaited_once() + + asyncio.run(_run()) + + +def test_evaluate_services_returns_requested_shape() -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + backend_calls: list[tuple[str, dict]] = [] + + async def _probe_agent_runtime() -> ResourceCheck: + return ResourceCheck( + status="fail", + message="down", + http_status_code=503, + http_status_desc="Service Unavailable", + ) + + async def _probe_agent_backend(*, agent_name: str) -> ResourceCheck: + backend_calls.append((agent_name, {})) + if agent_name == "conta": + return ResourceCheck( + status="ok", + http_status_code=204, + http_status_desc="No Content", + ) + return ResourceCheck( + status="ok", + http_status_code=200, + http_status_desc="oferta healthy", + ) + + async def _probe_stt() -> ResourceCheck: + return ResourceCheck( + status="ok", + http_status_code=200, + http_status_desc="sofya ready", + ) + + async def _probe_tts() -> ResourceCheck: + return ResourceCheck(status="ok", message="xAI TTS websocket connected") + + with ( + mock.patch.object(manager, "probe_agent_runtime", side_effect=_probe_agent_runtime), + mock.patch.object(manager, "probe_configured_agent_backend_service", side_effect=_probe_agent_backend), + mock.patch.object(manager, "probe_configured_stt_service", side_effect=_probe_stt), + mock.patch.object(manager, "probe_configured_xai_tts_service", side_effect=_probe_tts), + ): + report = await manager.evaluate_services() + + assert report.is_healthy() is False + assert report.as_dict() == { + "status": "fail", + "checks": { + "agent_runtime": { + "htp_cod_status": 503, + "hhtp_cod_desc": "Service Unavailable", + }, + "agent_backend": { + "contas": { + "htp_cod_status": 204, + "hhtp_cod_desc": "No Content", + }, + "oferta": { + "htp_cod_status": 200, + "hhtp_cod_desc": "oferta healthy", + }, + }, + "stt": { + "sofya": { + "htp_cod_status": 200, + "hhtp_cod_desc": "sofya ready", + }, + }, + "tts": { + "xAI": { + "htp_cod_status": 200, + "hhtp_cod_desc": "xAI TTS websocket connected", + }, + }, + }, + } + assert backend_calls == [("conta", {}), ("oferta", {})] + + asyncio.run(_run()) + + +def test_configured_agent_backend_service_uses_real_env_urls_when_backend_is_fake(monkeypatch) -> None: + async def _run() -> None: + manager = BridgeReadinessManager() + monkeypatch.setenv("AGENT_BACKEND", "remote_ws_fake") + monkeypatch.setenv("REMOTE_AGENT_HEALTH_URL_CONTA", "http://contas.example/health") + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL_OFERTA", raising=False) + monkeypatch.delenv("REMOTE_AGENT_HEALTH_URL_OFERTAS", raising=False) + monkeypatch.setenv("REMOTE_AGENT_SSE_URL_OFERTA", "https://oferta.example/agent/execute") + monkeypatch.setenv("REMOTE_AGENT_SSE_TLS_VERIFY_OFERTA", "0") + monkeypatch.delenv("REMOTE_AGENT_SSE_TLS_VERIFY", raising=False) + calls: list[tuple[str, str, bool]] = [] + + async def _fake_probe(check_name: str, url: str, *, verify: bool = True) -> ResourceCheck: + calls.append((check_name, url, verify)) + return ResourceCheck(status="ok", http_status_code=200, http_status_desc="OK") + + with mock.patch.object(manager, "_probe_http_health", side_effect=_fake_probe): + contas = await manager.probe_configured_agent_backend_service(agent_name="conta") + oferta = await manager.probe_configured_agent_backend_service(agent_name="oferta") + + assert contas.status == "ok" + assert oferta.status == "ok" + assert calls == [ + ("agent_backend", "http://contas.example/health", True), + ("agent_backend", "https://oferta.example/health", False), + ] + + asyncio.run(_run()) + + +def test_health_resources_route_returns_503(monkeypatch) -> None: + report = ReadinessReport( + checks={ + "agent_runtime": ResourceCheck(status="fail", message="down"), + "agent_backend": ResourceCheck(status="ok"), + "stt": ResourceCheck(status="ok"), + "tts": ResourceCheck(status="ok"), + }, + failed_resources=["agent_runtime"], + active_connections=2, + max_connections=5, + cached=False, + ) + fake_manager = _FakeReadinessManager(report=report) + monkeypatch.setattr(main_module, "READINESS_GATE", fake_manager) + + with TestClient(main_module.app) as client: + response = client.get("/health/resources") + + assert response.status_code == 503 + assert response.json()["failed_resources"] == ["agent_runtime"] + assert fake_manager.default_calls == 1 + + +def test_health_services_route_returns_503(monkeypatch) -> None: + services_report = ServicesHealthReport( + checks={ + "agent_runtime": { + "htp_cod_status": 500, + "hhtp_cod_desc": "FAIL", + }, + "agent_backend": { + "contas": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + "oferta": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + }, + "stt": { + "sofya": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + }, + "tts": { + "xAI": { + "htp_cod_status": 200, + "hhtp_cod_desc": "SUCCESS", + }, + }, + }, + healthy=False, + ) + fake_manager = _FakeReadinessManager() + fake_manager.services_report = services_report + monkeypatch.setattr(main_module, "READINESS_GATE", fake_manager) + + with TestClient(main_module.app) as client: + response = client.get("/health/services") + + assert response.status_code == 503 + assert response.json() == services_report.as_dict() + assert fake_manager.services_calls == 1 + + +def test_ws_agent_rejects_on_capacity_with_stop_data(monkeypatch) -> None: + lease = _FakeLease() + fake_manager = _FakeReadinessManager( + reservation=CapacityReservation( + allowed=False, + lease=lease, + active_connections=1, + max_connections=1, + ) + ) + monkeypatch.setattr(main_module, "READINESS_GATE", fake_manager) + monkeypatch.setenv("LIVEKIT_URL", "ws://127.0.0.1:7880") + + with TestClient(main_module.app) as client: + with client.websocket_connect("/ws/agent") as websocket: + websocket.send_text(json.dumps(_build_start_payload())) + payload = websocket.receive_json() + + assert payload == build_capacity_stop_message() + assert fake_manager.evaluate_calls == [] + assert lease.release_calls == 0 + + +def test_ws_agent_rejects_on_readiness_failure_and_releases_capacity(monkeypatch) -> None: + lease = _FakeLease() + report = ReadinessReport( + checks={ + "agent_runtime": ResourceCheck(status="ok"), + "agent_backend": ResourceCheck(status="ok"), + "stt": ResourceCheck(status="fail", message="down"), + "tts": ResourceCheck(status="ok"), + }, + failed_resources=["stt"], + active_connections=1, + max_connections=2, + cached=False, + ) + fake_manager = _FakeReadinessManager( + reservation=CapacityReservation( + allowed=True, + lease=lease, + active_connections=1, + max_connections=2, + ), + report=report, + ) + monkeypatch.setattr(main_module, "READINESS_GATE", fake_manager) + monkeypatch.setenv("LIVEKIT_URL", "ws://127.0.0.1:7880") + + with TestClient(main_module.app) as client: + with client.websocket_connect("/ws/agent") as websocket: + websocket.send_text(json.dumps(_build_start_payload())) + ready_payload = websocket.receive_json() + while True: + message = websocket.receive() + if message.get("text") is not None: + payload = json.loads(message["text"]) + break + assert ready_payload["type"] == "ready" + + assert payload == { + "type": "stop", + "data": { + "status": "stop_stt_unavailable", + "reason": "resource_unhealthy", + "resource": "stt", + "failed_resources": ["stt"], + "phase": "in_session", + }, + } + assert fake_manager.evaluate_calls == [("conta", {})] + assert lease.release_calls == 1 + + +def test_ws_agent_sends_terminal_stop_on_in_session_bridge_failure(monkeypatch) -> None: + lease = _FakeLease() + fake_manager = _FakeReadinessManager( + reservation=CapacityReservation( + allowed=True, + lease=lease, + active_connections=1, + max_connections=2, + ) + ) + timeline = _FakeTimeline() + monkeypatch.setattr(main_module, "READINESS_GATE", fake_manager) + monkeypatch.setenv("LIVEKIT_URL", "ws://127.0.0.1:7880") + + class _FakeRoom: + def __init__(self) -> None: + self.handlers = {} + self.disconnected = 0 + + def on(self, event_name: str): + def _decorator(fn): + self.handlers[event_name] = fn + return fn + + return _decorator + + async def disconnect(self) -> None: + self.disconnected += 1 + + class _FakeAudioSource: + def __init__(self, *args, **kwargs) -> None: + self.args = args + self.kwargs = kwargs + + class _FakeLocalAudioTrack: + @staticmethod + def create_audio_track(name: str, source: object) -> object: + return {"name": name, "source": source} + + class _FakeTrackPublishOptions: + def __init__(self) -> None: + self.source = None + + class _FakeTrackSource: + SOURCE_MICROPHONE = "microphone" + + async def _pending(*args, **kwargs) -> None: + await asyncio.Event().wait() + + async def _raise_dispatch(*args, **kwargs) -> None: + raise RuntimeError("bridge exploded") + + finalize_calls: list[bool] = [] + + class _FakeRecorder: + async def finalize(self) -> None: + finalize_calls.append(True) + + bootstrap = SimpleNamespace( + call_id_ged="GED-001", + call_config={}, + remote_agent_context={"agent": "conta"}, + room_name="dev-room-123", + identity="ws-bridge-dev-001", + token="token-001", + protocol="PRT-001", + phone_number="5511999999999", + timeline=timeline, + dispatch_metadata={}, + ) + + monkeypatch.setattr(main_module, "build_bridge_session_bootstrap", lambda **kwargs: bootstrap) + monkeypatch.setattr( + main_module, + "create_entire_call_recorder_from_env", + lambda **kwargs: _FakeRecorder(), + ) + monkeypatch.setattr(main_module, "dispatch_agent", _raise_dispatch) + monkeypatch.setattr(main_module, "connect_publish_livekit", _pending) + monkeypatch.setattr(main_module, "publish_queue_to_livekit", _pending) + monkeypatch.setattr(main_module, "ws_out_loop", _pending) + monkeypatch.setattr(main_module, "notify_client_audio_enabled", _pending) + monkeypatch.setattr(main_module.rtc, "Room", _FakeRoom) + monkeypatch.setattr(main_module.rtc, "AudioSource", _FakeAudioSource) + monkeypatch.setattr(main_module.rtc, "LocalAudioTrack", _FakeLocalAudioTrack) + monkeypatch.setattr(main_module.rtc, "TrackPublishOptions", _FakeTrackPublishOptions) + monkeypatch.setattr(main_module.rtc, "TrackSource", _FakeTrackSource) + + with TestClient(main_module.app) as client: + with client.websocket_connect("/ws/agent") as websocket: + websocket.send_text(json.dumps(_build_start_payload())) + ready_payload = websocket.receive_json() + stop_payload = websocket.receive_json() + + assert ready_payload["type"] == "ready" + assert stop_payload == build_bridge_failed_stop_message() + assert fake_manager.evaluate_calls == [("conta", {})] + assert lease.release_calls == 1 + assert finalize_calls == [True] diff --git a/tests/ws_gateway/test_session_audio.py b/tests/ws_gateway/test_session_audio.py new file mode 100644 index 0000000..fb27673 --- /dev/null +++ b/tests/ws_gateway/test_session_audio.py @@ -0,0 +1,290 @@ +from __future__ import annotations + +import asyncio +import json +import logging +import time +from types import SimpleNamespace + +from app.ws_gateway.session_audio import ( + enable_client_audio, + notify_client_audio_enabled, + stream_agent_audio_when_ready, + watch_mock_stop_after_first_audio, +) +from app.ws_gateway.session_lifecycle import RoomLifecycleState + + +class _FakeTimeline: + def __init__(self) -> None: + self.events: list[tuple[str, dict]] = [] + + def emit(self, event: str, **fields) -> None: + self.events.append((event, fields)) + + +class _FakeLocalParticipant: + def __init__(self) -> None: + self.published: list[dict] = [] + + async def publish_data(self, payload: str, **kwargs) -> None: + self.published.append({"payload": json.loads(payload), "kwargs": kwargs}) + + +class _FakeRoom: + def __init__(self) -> None: + self.local_participant = _FakeLocalParticipant() + + +class _FakeParticipant: + def __init__(self, identity: str) -> None: + self.identity = identity + + +class _FakeWebSocket: + def __init__(self) -> None: + self.text_messages: list[dict] = [] + self.closed = False + + async def send_text(self, payload: str) -> None: + self.text_messages.append(json.loads(payload)) + + async def close(self) -> None: + self.closed = True + + +def test_enable_client_audio_sets_gate_and_emits_timeline() -> None: + async def _run() -> None: + timeline = _FakeTimeline() + logger = logging.getLogger("test.session_audio.enable") + client_audio_enabled = asyncio.Event() + + await enable_client_audio( + logger=logger, + timeline=timeline, + client_audio_enabled=client_audio_enabled, + room_name="room-dev-123", + protocol="PRT-1", + ) + + assert client_audio_enabled.is_set() + assert timeline.events == [ + ("client_audio_enabled", {}), + ] + + asyncio.run(_run()) + + +def test_enable_client_audio_is_idempotent() -> None: + async def _run() -> None: + timeline = _FakeTimeline() + logger = logging.getLogger("test.session_audio.enable.idempotent") + client_audio_enabled = asyncio.Event() + + await enable_client_audio( + logger=logger, + timeline=timeline, + client_audio_enabled=client_audio_enabled, + room_name="room-dev-123", + protocol="PRT-1", + ) + await enable_client_audio( + logger=logger, + timeline=timeline, + client_audio_enabled=client_audio_enabled, + room_name="room-dev-123", + protocol="PRT-1", + ) + + assert client_audio_enabled.is_set() + assert timeline.events == [("client_audio_enabled", {})] + + asyncio.run(_run()) + + +def test_notify_client_audio_enabled_publishes_bridge_control_message() -> None: + async def _run() -> None: + timeline = _FakeTimeline() + logger = logging.getLogger("test.session_audio.notify") + room = _FakeRoom() + track_published = asyncio.Event() + track_published.set() + agent_ready = asyncio.Event() + agent_ready.set() + lifecycle = RoomLifecycleState(agent_participant=_FakeParticipant("agent-001")) + + await notify_client_audio_enabled( + room=room, + logger=logger, + timeline=timeline, + room_name="room-dev-123", + protocol="PRT-1", + track_published=track_published, + agent_ready=agent_ready, + lifecycle=lifecycle, + ) + + assert room.local_participant.published == [ + { + "payload": { + "type": "client_audio_enabled", + "room": "room-dev-123", + "protocol": "PRT-1", + }, + "kwargs": { + "reliable": True, + "destination_identities": ["agent-001"], + "topic": "bridge.control", + }, + } + ] + assert timeline.events == [ + ( + "bridge_control_sent", + {"control_type": "client_audio_enabled", "agent_identity": "agent-001"}, + ) + ] + + asyncio.run(_run()) + + +def test_stream_agent_audio_when_ready_uses_lifecycle_participant() -> None: + async def _run() -> None: + agent_ready = asyncio.Event() + agent_ready.set() + participant = _FakeParticipant("agent-001") + lifecycle = RoomLifecycleState(agent_participant=participant) + calls: list[dict] = [] + + async def _stream_agent_audio(agent_q, agent_participant, activity, timeline=None) -> None: + calls.append( + { + "agent_q": agent_q, + "agent_participant": agent_participant, + "activity": activity, + "timeline": timeline, + } + ) + + queue: asyncio.Queue[bytes] = asyncio.Queue() + activity = SimpleNamespace(last_client_in=0.0, last_agent_out=0.0) + timeline = _FakeTimeline() + + await stream_agent_audio_when_ready( + agent_ready=agent_ready, + lifecycle=lifecycle, + agent_q=queue, + activity=activity, + timeline=timeline, + stream_agent_audio=_stream_agent_audio, + ) + + assert calls == [ + { + "agent_q": queue, + "agent_participant": participant, + "activity": activity, + "timeline": timeline, + } + ] + + asyncio.run(_run()) + + +def test_stream_agent_audio_when_ready_repeats_for_reconnected_agent() -> None: + async def _run() -> None: + agent_ready = asyncio.Event() + agent_ready.set() + first = _FakeParticipant("agent-001") + second = _FakeParticipant("agent-002") + lifecycle = RoomLifecycleState(agent_participant=first) + calls: list[str] = [] + second_call = asyncio.Event() + + async def _stream_agent_audio(agent_q, agent_participant, activity, timeline=None) -> None: + calls.append(agent_participant.identity) + if agent_participant is first: + lifecycle.agent_participant = None + lifecycle.agent_generation += 1 + lifecycle.agent_connected.clear() + + async def _reconnect() -> None: + await asyncio.sleep(0.01) + lifecycle.agent_participant = second + lifecycle.agent_generation += 1 + lifecycle.agent_connected.set() + + asyncio.create_task(_reconnect()) + return + + second_call.set() + await asyncio.Event().wait() + + queue: asyncio.Queue[bytes] = asyncio.Queue() + activity = SimpleNamespace(last_client_in=0.0, last_agent_out=0.0) + task = asyncio.create_task( + stream_agent_audio_when_ready( + agent_ready=agent_ready, + lifecycle=lifecycle, + agent_q=queue, + activity=activity, + timeline=None, + stream_agent_audio=_stream_agent_audio, + repeat=True, + ) + ) + + await asyncio.wait_for(second_call.wait(), timeout=1.0) + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + assert calls == ["agent-001", "agent-002"] + + asyncio.run(_run()) + + +def test_watch_mock_stop_after_first_audio_sends_resolved_terminal_stop() -> None: + async def _run() -> None: + timeline = _FakeTimeline() + logger = logging.getLogger("test.session_audio.mock_stop") + ws = _FakeWebSocket() + first_agent_audio_sent = asyncio.Event() + activity = SimpleNamespace(last_agent_out=time.monotonic()) + agent_q: asyncio.Queue[bytes] = asyncio.Queue() + + task = asyncio.create_task( + watch_mock_stop_after_first_audio( + ws=ws, + logger=logger, + timeline=timeline, + room_name="room-dev-123", + protocol="PRT-1", + first_agent_audio_sent=first_agent_audio_sent, + activity=activity, + agent_q=agent_q, + silence_s=0.01, + reason="stage_done", + ) + ) + + first_agent_audio_sent.set() + await asyncio.sleep(0.03) + await task + + assert ws.text_messages == [ + { + "type": "stop", + "data": { + "status": "stop_resolvido_e_finalizado", + "reason": "stage_done", + "phase": "in_session", + }, + } + ] + assert ws.closed is True + assert timeline.events == [ + ("mock_stop_after_first_audio_armed", {"silence_ms": 10, "reason": "stage_done"}), + ("mock_stop_after_first_audio_sent", {"reason": "stage_done"}), + ] + + asyncio.run(_run()) diff --git a/tests/ws_gateway/test_session_bootstrap.py b/tests/ws_gateway/test_session_bootstrap.py new file mode 100644 index 0000000..f6407bd --- /dev/null +++ b/tests/ws_gateway/test_session_bootstrap.py @@ -0,0 +1,201 @@ +from __future__ import annotations + +import logging + +from app.ws_gateway.session_bootstrap import build_bridge_session_bootstrap +from app.ws_gateway.session_start import StartSessionContext + + +def _build_start_ctx(**data_overrides: str) -> StartSessionContext: + data = { + "agent": "oferta", + "ani": "5511888888888", + "gsm": "5511999999999", + "session_id": "550e8400-e29b-41d4-a716-446655440100", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-1234567890", + **data_overrides, + } + payload = { + "type": "start", + "data": data, + "callConfig": { + "tts": {"provider": "azure"}, + }, + } + return StartSessionContext( + payload=payload, + data=data, + agent_data=dict(data.get("agentData") or {}), + audio_format={}, + call_config=payload["callConfig"], + session_data={"gsm": data.get("gsm", ""), "msisdn": data.get("gsm", ""), **data}, + intro="intro", + nudge="nudge", + ) + + +def test_build_bridge_session_bootstrap_uses_protocol_from_data_and_builds_dispatch_metadata( + monkeypatch, +) -> None: + logger = logging.getLogger("test.session_bootstrap") + monkeypatch.setenv("CALL_TIMELINE_DIR", "/tmp/timeline-bootstrap") + monkeypatch.setenv("CALL_TIMELINE_ENABLED", "0") + + data = { + "agent": "conta", + "ani": "5511888888888", + "gsm": "5511888888888", + "session_id": "550e8400-e29b-41d4-a716-446655440101", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-321", + "callIdGed": "GED-321", + "protocolo": "PRT-321", + "agentData": { + "idFatura": "fat-999", + }, + } + start_ctx = StartSessionContext( + payload={ + "type": "start", + "data": data, + "callConfig": { + "tts": {"provider": "azure"}, + }, + }, + data=data, + agent_data={"idFatura": "fat-999"}, + audio_format={}, + call_config={"tts": {"provider": "azure"}}, + session_data={"gsm": "5511888888888", "msisdn": "5511888888888", **data}, + intro="intro", + nudge="nudge", + ) + + uuids = iter( + [ + type("U", (), {"hex": "room123456789"})(), + type("U", (), {"hex": "ident12345678"})(), + ] + ) + token_calls: list[tuple[str, str]] = [] + + bootstrap = build_bridge_session_bootstrap( + start_ctx=start_ctx, + app_env="dev", + livekit_room="custom-room", + default_protocol="fallback-proto", + token_factory=lambda identity, room_name: token_calls.append((identity, room_name)) or "jwt-token", + logger=logger, + now_fn=lambda: 1234567890, + uuid_factory=lambda: next(uuids), + ) + + assert bootstrap.call_id_ged == "GED-321" + assert bootstrap.room_name == "custom-room-room1234" + assert bootstrap.identity == "ws-bridge-dev-ident123" + assert bootstrap.token == "jwt-token" + assert bootstrap.protocol == "PRT-321" + assert bootstrap.phone_number == "5511888888888" + assert bootstrap.remote_agent_context["agent"] == "conta" + assert bootstrap.remote_agent_context["ID_FATURA"] == "fat-999" + assert bootstrap.remote_agent_context["callIdGed"] == "GED-321" + assert bootstrap.call_config["tts"]["provider"] == "azure" + assert bootstrap.dispatch_metadata["bridge_identity"] == "ws-bridge-dev-ident123" + assert bootstrap.dispatch_metadata["protocol"] == "PRT-321" + assert bootstrap.dispatch_metadata["call_id_ged"] == "GED-321" + assert bootstrap.dispatch_metadata["agent_starts_conversation"] is True + assert bootstrap.dispatch_metadata["remote_agent"]["ID_FATURA"] == "fat-999" + assert bootstrap.dispatch_metadata["timeline_id"] == "custom-room-room1234" + assert bootstrap.dispatch_metadata["session_id"] == "550e8400-e29b-41d4-a716-446655440101" + assert token_calls == [("ws-bridge-dev-ident123", "custom-room-room1234")] + + +def test_build_bridge_session_bootstrap_falls_back_to_router_call_key_when_protocol_is_missing( + monkeypatch, +) -> None: + logger = logging.getLogger("test.session_bootstrap") + monkeypatch.setenv("CALL_TIMELINE_DIR", "/tmp/timeline-bootstrap") + monkeypatch.setenv("CALL_TIMELINE_ENABLED", "0") + + start_ctx = _build_start_ctx() + + bootstrap = build_bridge_session_bootstrap( + start_ctx=start_ctx, + app_env="dev", + livekit_room="custom-room", + default_protocol="fallback-proto", + token_factory=lambda *_args: "jwt-token", + logger=logger, + now_fn=lambda: 1234567890, + uuid_factory=lambda: type("U", (), {"hex": "room123456789"})(), + ) + + assert bootstrap.protocol == "RCK-001" + assert bootstrap.dispatch_metadata["protocol"] == "RCK-001" + + +def test_build_bridge_session_bootstrap_falls_back_to_generated_protocol_and_phone(monkeypatch) -> None: + logger = logging.getLogger("test.session_bootstrap") + monkeypatch.setenv("CALL_TIMELINE_DIR", "/tmp/timeline-bootstrap") + monkeypatch.setenv("CALL_TIMELINE_ENABLED", "0") + + start_ctx = _build_start_ctx(routerCallKey="", gsm="") + start_ctx = StartSessionContext( + payload=start_ctx.payload, + data=start_ctx.data, + agent_data=start_ctx.agent_data, + audio_format=start_ctx.audio_format, + call_config=start_ctx.call_config, + session_data={**start_ctx.session_data, "phone": "5511777777777"}, + intro=start_ctx.intro, + nudge=start_ctx.nudge, + ) + uuids = iter( + [ + type("U", (), {"hex": "roomabcdef1234"})(), + type("U", (), {"hex": "identfedcba98"})(), + ] + ) + + bootstrap = build_bridge_session_bootstrap( + start_ctx=start_ctx, + app_env="qa", + livekit_room="", + default_protocol="", + token_factory=lambda *_args: "jwt-token", + logger=logger, + now_fn=lambda: 1712345678, + uuid_factory=lambda: next(uuids), + ) + + assert bootstrap.room_name == "qa-room-roomabcd" + assert bootstrap.identity == "ws-bridge-qa-identfed" + assert bootstrap.protocol == "WS-GED-1234567890-1712345678" + assert bootstrap.phone_number == "5511777777777" + assert bootstrap.dispatch_metadata["session_id"] == "550e8400-e29b-41d4-a716-446655440100" + assert bootstrap.dispatch_metadata["agent_starts_conversation"] is True + assert bootstrap.dispatch_metadata["session_data"]["phone"] == "5511777777777" + + +def test_build_bridge_session_bootstrap_uses_explicit_session_id_without_protocol_fallback(monkeypatch) -> None: + logger = logging.getLogger("test.session_bootstrap.session_id") + monkeypatch.setenv("CALL_TIMELINE_DIR", "/tmp/timeline-bootstrap") + monkeypatch.setenv("CALL_TIMELINE_ENABLED", "0") + + start_ctx = _build_start_ctx(protocolo="PRT-001", session_id="sess-001") + + bootstrap = build_bridge_session_bootstrap( + start_ctx=start_ctx, + app_env="dev", + livekit_room="custom-room", + default_protocol="fallback-proto", + token_factory=lambda *_args: "jwt-token", + logger=logger, + now_fn=lambda: 1234567890, + uuid_factory=lambda: type("U", (), {"hex": "room123456789"})(), + ) + + assert bootstrap.protocol == "PRT-001" + assert bootstrap.dispatch_metadata["session_id"] == "sess-001" diff --git a/tests/ws_gateway/test_session_lifecycle.py b/tests/ws_gateway/test_session_lifecycle.py new file mode 100644 index 0000000..774e96a --- /dev/null +++ b/tests/ws_gateway/test_session_lifecycle.py @@ -0,0 +1,365 @@ +from __future__ import annotations + +import asyncio +import json +import logging + +from app.ws_gateway.session_lifecycle import ( + RoomLifecycleState, + register_room_lifecycle_handlers, + watch_call_done, +) + + +class _FakeRoom: + def __init__(self) -> None: + self.handlers = {} + + def on(self, event_name: str): + def _decorator(fn): + self.handlers[event_name] = fn + return fn + + return _decorator + + +class _FakeTimeline: + def __init__(self) -> None: + self.events: list[tuple[str, dict]] = [] + + def emit(self, event: str, **fields) -> None: + self.events.append((event, fields)) + + +class _FakeParticipant: + def __init__(self, identity: str) -> None: + self.identity = identity + + +class _FakePacket: + def __init__(self, topic: str, data) -> None: + self.topic = topic + self.data = data + + +class _FakeWebSocket: + def __init__(self) -> None: + self.sent: list[dict] = [] + self.closed = 0 + self.close_calls: list[dict] = [] + + async def send_text(self, payload: str) -> None: + self.sent.append(json.loads(payload)) + + async def close(self, code: int = 1000, reason: str = "") -> None: + self.closed += 1 + self.close_calls.append({"code": code, "reason": reason}) + + +def test_register_room_lifecycle_handlers_updates_state_and_emits_timeline() -> None: + room = _FakeRoom() + timeline = _FakeTimeline() + logger = logging.getLogger("test.session_lifecycle.register") + call_done = asyncio.Event() + agent_ready = asyncio.Event() + state = RoomLifecycleState() + + register_room_lifecycle_handlers( + room=room, + logger=logger, + timeline=timeline, + room_name="room-dev-123", + protocol="PRT-1", + dispatch_started_at=100.0, + agent_ready=agent_ready, + call_done=call_done, + state=state, + is_agent=lambda participant: str(getattr(participant, "identity", "")).startswith("agent"), + monotonic_fn=lambda: 100.123, + ) + + participant = _FakeParticipant("agent-001") + room.handlers["participant_connected"](participant) + room.handlers["data_received"]( + _FakePacket("agent.stage", b'{"stage":"DONE","reason":"completed"}') + ) + room.handlers["participant_disconnected"](participant) + + assert agent_ready.is_set() + assert call_done.is_set() + assert state.agent_participant is None + assert state.done_payload == {"stage": "DONE", "reason": "completed"} + assert timeline.events == [ + ("agent_join", {"agent_identity": "agent-001", "dispatch_dt_ms": 123}), + ( + "done_packet_received", + {"reason": "completed", "payload": {"stage": "DONE", "reason": "completed"}}, + ), + ("agent_leave", {"agent_identity": "agent-001"}), + ] + + +def test_agent_disconnect_without_done_notifies_recovery_handler() -> None: + async def _run() -> None: + room = _FakeRoom() + timeline = _FakeTimeline() + logger = logging.getLogger("test.session_lifecycle.agent_disconnect") + call_done = asyncio.Event() + agent_ready = asyncio.Event() + state = RoomLifecycleState() + calls: list[dict] = [] + + async def _on_agent_disconnect(agent_identity: str, generation: int) -> None: + calls.append({"agent_identity": agent_identity, "generation": generation}) + + register_room_lifecycle_handlers( + room=room, + logger=logger, + timeline=timeline, + room_name="room-dev-123", + protocol="PRT-1", + dispatch_started_at=100.0, + agent_ready=agent_ready, + call_done=call_done, + state=state, + is_agent=lambda participant: str(getattr(participant, "identity", "")).startswith("agent"), + on_agent_disconnect=_on_agent_disconnect, + monotonic_fn=lambda: 100.123, + ) + + participant = _FakeParticipant("agent-001") + room.handlers["participant_connected"](participant) + connected_generation = state.agent_generation + room.handlers["participant_disconnected"](participant) + await asyncio.sleep(0) + + assert state.agent_participant is None + assert not state.agent_connected.is_set() + assert state.agent_generation == connected_generation + 1 + assert calls == [ + {"agent_identity": "agent-001", "generation": connected_generation + 1} + ] + + asyncio.run(_run()) + + +def test_watch_call_done_sends_stop_and_closes_socket() -> None: + async def _run() -> None: + timeline = _FakeTimeline() + ws = _FakeWebSocket() + logger = logging.getLogger("test.session_lifecycle.done") + call_done = asyncio.Event() + state = RoomLifecycleState(done_payload={"reason": "finished"}) + call_done.set() + + await watch_call_done( + call_done=call_done, + state=state, + ws=ws, + logger=logger, + room_name="room-dev-123", + protocol="PRT-1", + timeline=timeline, + ) + + assert ws.sent == [ + { + "type": "stop", + "data": { + "status": "stop_resolvido_e_finalizado", + "reason": "finished", + "phase": "in_session", + }, + } + ] + assert ws.closed == 1 + assert ws.close_calls == [{"code": 1000, "reason": "call_done:finished"}] + assert timeline.events == [ + ("call_done", {"reason": "finished"}), + ("stop_sent", {"reason": "finished"}), + ] + + asyncio.run(_run()) + + +def test_agent_debug_packet_is_forwarded_when_enabled() -> None: + async def _run() -> None: + room = _FakeRoom() + ws = _FakeWebSocket() + register_room_lifecycle_handlers( + room=room, + logger=logging.getLogger("test.session_lifecycle.debug"), + timeline=_FakeTimeline(), + room_name="room-load-1", + protocol="LOAD-1", + dispatch_started_at=1.0, + agent_ready=asyncio.Event(), + call_done=asyncio.Event(), + state=RoomLifecycleState(), + is_agent=lambda _participant: False, + ws=ws, + debug_events_enabled=True, + ) + event = { + "type": "debug_event", + "version": 1, + "source": "agent", + "event": "agent.speech.finished", + "data": {"text": "Resposta", "stage": "ARGUMENTATION"}, + } + room.handlers["data_received"]( + _FakePacket("agent.debug", json.dumps(event).encode("utf-8")) + ) + await asyncio.sleep(0) + + assert ws.sent == [{**event, "stress_test": True}] + + asyncio.run(_run()) + + +def test_agent_debug_packet_is_not_forwarded_when_disabled() -> None: + async def _run() -> None: + room = _FakeRoom() + ws = _FakeWebSocket() + register_room_lifecycle_handlers( + room=room, + logger=logging.getLogger("test.session_lifecycle.debug.disabled"), + timeline=_FakeTimeline(), + room_name="room-regular-1", + protocol="REGULAR-1", + dispatch_started_at=1.0, + agent_ready=asyncio.Event(), + call_done=asyncio.Event(), + state=RoomLifecycleState(), + is_agent=lambda _participant: False, + ws=ws, + debug_events_enabled=False, + ) + event = { + "type": "debug_event", + "version": 1, + "source": "agent", + "event": "stt.completed", + "data": {"duration_ms": 100}, + } + room.handlers["data_received"]( + _FakePacket("agent.debug", json.dumps(event).encode("utf-8")) + ) + await asyncio.sleep(0) + + assert ws.sent == [] + + asyncio.run(_run()) + + +def test_watch_call_done_maps_no_user_response_to_long_silence_status() -> None: + async def _run() -> None: + timeline = _FakeTimeline() + ws = _FakeWebSocket() + logger = logging.getLogger("test.session_lifecycle.done.long_silence") + call_done = asyncio.Event() + state = RoomLifecycleState(done_payload={"reason": "no_user_response"}) + call_done.set() + + await watch_call_done( + call_done=call_done, + state=state, + ws=ws, + logger=logger, + room_name="room-dev-123", + protocol="PRT-1", + timeline=timeline, + ) + + assert ws.sent == [ + { + "type": "stop", + "data": { + "status": "stop_silencio_longo", + "reason": "no_user_response", + "phase": "in_session", + }, + } + ] + assert ws.closed == 1 + + asyncio.run(_run()) + + +def test_watch_call_done_uses_default_status_when_reason_is_not_exact_match() -> None: + async def _run() -> None: + timeline = _FakeTimeline() + ws = _FakeWebSocket() + logger = logging.getLogger("test.session_lifecycle.done.default") + call_done = asyncio.Event() + state = RoomLifecycleState(done_payload={"reason": "finished"}) + call_done.set() + + await watch_call_done( + call_done=call_done, + state=state, + ws=ws, + logger=logger, + room_name="room-dev-123", + protocol="PRT-1", + timeline=timeline, + ) + + assert ws.sent == [ + { + "type": "stop", + "data": { + "status": "stop_resolvido_e_finalizado", + "reason": "finished", + "phase": "in_session", + }, + } + ] + assert ws.closed == 1 + + asyncio.run(_run()) + + +def test_watch_call_done_preserves_terminal_status_from_agent_payload() -> None: + async def _run() -> None: + timeline = _FakeTimeline() + ws = _FakeWebSocket() + logger = logging.getLogger("test.session_lifecycle.done.custom") + call_done = asyncio.Event() + state = RoomLifecycleState( + done_payload={ + "stage": "DONE", + "status": "stop_agent_backend_unavailable", + "reason": "resource_unhealthy", + "resource": "agent_backend", + "failed_resources": ["agent_backend"], + "phase": "in_session", + } + ) + call_done.set() + + await watch_call_done( + call_done=call_done, + state=state, + ws=ws, + logger=logger, + room_name="room-dev-123", + protocol="PRT-1", + timeline=timeline, + ) + + assert ws.sent == [ + { + "type": "stop", + "data": { + "status": "stop_agent_backend_unavailable", + "reason": "resource_unhealthy", + "resource": "agent_backend", + "failed_resources": ["agent_backend"], + "phase": "in_session", + }, + } + ] + assert ws.closed == 1 + + asyncio.run(_run()) diff --git a/tests/ws_gateway/test_session_start.py b/tests/ws_gateway/test_session_start.py new file mode 100644 index 0000000..252b0f4 --- /dev/null +++ b/tests/ws_gateway/test_session_start.py @@ -0,0 +1,385 @@ +from __future__ import annotations + +import asyncio +import json + +import pytest + +from app.ws_gateway.session_start import ( + build_remote_agent_context, + parse_start_payload, + parse_transferencia_session_id_payload, + recv_start_message, +) + + +def test_parse_start_payload_builds_context_from_data() -> None: + ctx = parse_start_payload( + { + "type": "start", + "data": { + "agent": "oferta", + "ani": "5511999999999", + "gsm": "5511999999999", + "session_id": "550e8400-e29b-41d4-a716-446655440000", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-001", + "protocolo": "PRT-001", + }, + "audioFormat": {"encoding": "linear16"}, + "callConfig": {"agentBackend": "remote_ws"}, + } + ) + + assert ctx.data["agent"] == "oferta" + assert ctx.data["session_id"] == "550e8400-e29b-41d4-a716-446655440000" + assert ctx.data["protocolo"] == "PRT-001" + assert ctx.agent_data == {} + assert ctx.audio_format == {"encoding": "linear16"} + assert ctx.call_config == {"agentBackend": "remote_ws"} + assert ctx.session_data["msisdn"] == "5511999999999" + assert ctx.session_data["audioFormat"] == {"encoding": "linear16"} + assert ctx.agent_starts_conversation is True + assert ctx.intro == "" + assert ctx.nudge == "Alô, você ainda está aí?" + + +def test_parse_start_payload_lets_conta_agent_own_first_message() -> None: + ctx = parse_start_payload( + { + "type": "start", + "data": { + "agent": "conta", + "ani": "5511999999999", + "gsm": "5511999999999", + "session_id": "550e8400-e29b-41d4-a716-446655440001", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-001", + "agentData": { + "idFatura": "fat-123", + }, + }, + } + ) + + assert ctx.agent_starts_conversation is True + assert ctx.intro == "" + assert ctx.nudge == "Alô, você ainda está aí?" + + +def test_parse_start_payload_rejects_non_start_message() -> None: + with pytest.raises(RuntimeError): + parse_start_payload({"type": "ping"}) + + +def test_parse_start_payload_rejects_invalid_fake_agent_responses() -> None: + with pytest.raises(RuntimeError, match="callConfig invalido"): + parse_start_payload( + { + "type": "start", + "data": { + "agent": "conta", + "ani": "5511999999999", + "gsm": "5511999999999", + "routerCallKeyDay": "20260819", + "routerCallKey": "RCK-1", + "callIdGed": "GED-1", + "session_id": "session-1", + "agentData": {"idFatura": "FAT-1"}, + }, + "callConfig": { + "agentBackend": "remote_ws_fake", + "agentFake": {"responses": "resposta curta;outra curta"}, + }, + } + ) + + +def test_parse_start_payload_requires_id_fatura_for_conta() -> None: + with pytest.raises(RuntimeError, match="idFatura"): + parse_start_payload( + { + "type": "start", + "data": { + "agent": "conta", + "ani": "5511999999999", + "gsm": "5511999999999", + "session_id": "550e8400-e29b-41d4-a716-446655440002", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-001", + }, + } + ) + + +def test_parse_start_payload_requires_protocolo_for_oferta() -> None: + with pytest.raises(RuntimeError, match="protocolo"): + parse_start_payload( + { + "type": "start", + "data": { + "agent": "oferta", + "ani": "5511999999999", + "gsm": "5511999999999", + "session_id": "550e8400-e29b-41d4-a716-446655440003", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-001", + }, + } + ) + + +def test_parse_start_payload_requires_session_id() -> None: + with pytest.raises(RuntimeError, match="session_id"): + parse_start_payload( + { + "type": "start", + "data": { + "agent": "oferta", + "ani": "5511999999999", + "gsm": "5511999999999", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-001", + "protocolo": "PRT-001", + }, + } + ) + + +def test_parse_start_payload_accepts_camel_session_id_and_normalizes() -> None: + ctx = parse_start_payload( + { + "type": "start", + "data": { + "agent": "oferta", + "ani": "5511999999999", + "gsm": "5511999999999", + "sessionId": "550e8400-e29b-41d4-a716-446655440004", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-001", + "protocolo": "PRT-001", + }, + } + ) + + assert ctx.data["session_id"] == "550e8400-e29b-41d4-a716-446655440004" + + +def test_parse_transferencia_session_id_payload_extracts_session_id() -> None: + session_id = parse_transferencia_session_id_payload( + { + "type": "transferencia_session_id", + "data": {"session_id": "550e8400-e29b-41d4-a716-446655440005"}, + } + ) + + assert session_id == "550e8400-e29b-41d4-a716-446655440005" + + +def test_recv_start_message_uses_transferencia_session_id_before_start() -> None: + class FakeWebSocket: + def __init__(self) -> None: + self.messages = [ + { + "type": "transferencia_session_id", + "data": {"session_id": "550e8400-e29b-41d4-a716-446655440006"}, + }, + { + "type": "start", + "data": { + "agent": "oferta", + "ani": "5511999999999", + "gsm": "5511999999999", + "routerCallKeyDay": "20260409", + "routerCallKey": "RCK-001", + "callIdGed": "GED-001", + "protocolo": "PRT-001", + }, + }, + ] + + async def receive_text(self) -> str: + return json.dumps(self.messages.pop(0)) + + ctx = asyncio.run(recv_start_message(FakeWebSocket())) + + assert ctx.data["session_id"] == "550e8400-e29b-41d4-a716-446655440006" + + +def test_build_remote_agent_context_uses_camel_case_payload() -> None: + context = build_remote_agent_context( + data={ + "agent": "contas", + "ani": "5511888888888", + "gsm": "5511777777777", + "routerCallKeyDay": "20260406", + "routerCallKey": "RCK-123", + "callIdGed": "GED-777", + }, + agent_data={"idFatura": "fat-123"}, + ) + + assert context == { + "agent": "conta", + "RouterCallKeyDay": "20260406", + "RouterCallKey": "RCK-123", + "ANI": "5511888888888", + "GSM": "5511777777777", + "msisdn": "5511777777777", + "callIdGed": "GED-777", + "ID_FATURA": "fat-123", + "current_invoice_number": "fat-123", + } + + +def test_build_remote_agent_context_for_conta_forwards_protocol_id() -> None: + context = build_remote_agent_context( + data={ + "agent": "conta", + "ani": "5511888888888", + "gsm": "5511777777777", + "routerCallKeyDay": "20260406", + "routerCallKey": "RCK-123", + "callIdGed": "GED-777", + "protocol_id": "PRT-777", + }, + agent_data={"idFatura": "fat-123"}, + ) + + assert context["protocol_id"] == "PRT-777" + assert context["protocolo"] == "PRT-777" + assert context["protocolNumber"] == "PRT-777" + + +def test_build_remote_agent_context_for_oferta_uses_protocolo_and_never_fatura() -> None: + context = build_remote_agent_context( + data={ + "agent": "ofertas", + "ani": "5511888888888", + "gsm": "5511777777777", + "routerCallKeyDay": "20260406", + "routerCallKey": "RCK-123", + "callIdGed": "GED-777", + "protocolo": "PRT-777", + "assetId": "asset-777", + "channelId": "ura", + }, + agent_data={"idFatura": "fat-should-not-leak"}, + ) + + assert context == { + "agent": "oferta", + "RouterCallKeyDay": "20260406", + "RouterCallKey": "RCK-123", + "ANI": "5511888888888", + "GSM": "5511777777777", + "msisdn": "5511777777777", + "callIdGed": "GED-777", + "protocolo": "PRT-777", + "protocolNumber": "PRT-777", + "channelId": "ura", + "assetId": "asset-777", + } + + +def test_build_remote_agent_context_only_forwards_explicit_session_id() -> None: + without_session = build_remote_agent_context( + data={ + "agent": "oferta", + "ani": "5511888888888", + "gsm": "5511777777777", + "routerCallKeyDay": "20260406", + "routerCallKey": "RCK-123", + "callIdGed": "GED-777", + "protocolo": "PRT-777", + }, + agent_data={}, + ) + with_session = build_remote_agent_context( + data={ + **without_session, + "agent": "oferta", + "routerCallKeyDay": "20260406", + "routerCallKey": "RCK-123", + "ani": "5511888888888", + "gsm": "5511777777777", + "sessionId": "sess-777", + }, + agent_data={}, + ) + + assert "sessionId" not in without_session + assert "session_id" not in without_session + assert with_session["sessionId"] == "sess-777" + assert with_session["session_id"] == "sess-777" + + +def test_build_remote_agent_context_forwards_explicit_message_id() -> None: + context = build_remote_agent_context( + data={ + "agent": "conta", + "ani": "5511888888888", + "gsm": "5511777777777", + "routerCallKeyDay": "20260406", + "routerCallKey": "RCK-123", + "callIdGed": "GED-777", + "messageId": "msg-777", + }, + agent_data={}, + ) + + assert context["message_id"] == "msg-777" + + +def test_parse_start_payload_does_not_enable_debug_events_from_client_flag_alone() -> None: + ctx = parse_start_payload( + { + "type": "start", + "debug": {"events": True}, + "data": { + "agent": "conta", + "ani": "5511999999999", + "gsm": "5511999999999", + "session_id": "load-debug-1", + "routerCallKeyDay": "20260817", + "routerCallKey": "LOAD-1", + "callIdGed": "GED-LOAD-1", + "agentData": {"idFatura": "fat-load-1"}, + }, + } + ) + + assert ctx.stress_test is False + assert ctx.debug_events_enabled is False + + +def test_parse_start_payload_enables_debug_events_for_scripted_fake_agent() -> None: + response = "Esta resposta simulada possui tamanho suficiente para validar o contrato." + ctx = parse_start_payload( + { + "type": "start", + "data": { + "agent": "conta", + "ani": "5511999999999", + "gsm": "5511999999999", + "session_id": "stress-debug-1", + "routerCallKeyDay": "20260817", + "routerCallKey": "STRESS-1", + "callIdGed": "GED-STRESS-1", + "agentData": {"idFatura": "fat-stress-1"}, + }, + "callConfig": { + "agentBackend": "remote_ws_fake", + "agentFake": {"responses": f"{response};{response}"}, + }, + } + ) + + assert ctx.stress_test is True + assert ctx.debug_events_enabled is True diff --git a/tests/ws_gateway/test_voice_client.py b/tests/ws_gateway/test_voice_client.py new file mode 100644 index 0000000..85b2fec --- /dev/null +++ b/tests/ws_gateway/test_voice_client.py @@ -0,0 +1,56 @@ +from pathlib import Path + + +VOICE_CLIENT_HTML = ( + Path(__file__).resolve().parents[2] + / "src" + / "app" + / "ws_gateway" + / "voice_client.html" +) + + +def test_voice_client_exposes_required_protocolo_for_oferta() -> None: + html = VOICE_CLIENT_HTML.read_text(encoding="utf-8") + + assert "oferta: [" in html + assert 'key: "protocolo"' in html + assert 'inputId: "agentFieldProtocolo"' in html + assert 'label: "PROTOCOLO"' in html + assert 'required: true' in html + assert 'target: "data"' in html + + +def test_voice_client_keeps_id_fatura_inside_agent_data() -> None: + html = VOICE_CLIENT_HTML.read_text(encoding="utf-8") + + assert "const agentData = {};" in html + assert 'key: "idFatura"' in html + assert 'data.agentData = agentData;' in html + + +def test_voice_client_exposes_required_protocol_id_for_conta() -> None: + html = VOICE_CLIENT_HTML.read_text(encoding="utf-8") + + assert "conta: [" in html + assert 'key: "protocol_id"' in html + assert 'inputId: "agentFieldProtocolId"' in html + assert 'label: "PROTOCOL_ID"' in html + assert 'target: "data"' in html + + +def test_voice_client_sends_session_id_in_start_data() -> None: + html = VOICE_CLIENT_HTML.read_text(encoding="utf-8") + + assert 'id="sessionId"' in html + assert "function defaultSessionId()" in html + assert "session_id: sessionId" in html + + +def test_voice_client_guards_stale_microphone_streams() -> None: + html = VOICE_CLIENT_HTML.read_text(encoding="utf-8") + + assert "function stopMicStreaming()" in html + assert "const wsAtStart = state.ws;" in html + assert "state.ws !== wsAtStart" in html + assert "mediaStream.getTracks().forEach((track) => track.stop());" in html