diff --git a/backend/chainlit/message.py b/backend/chainlit/message.py index 0700f630a7..c8d5e41079 100644 --- a/backend/chainlit/message.py +++ b/backend/chainlit/message.py @@ -58,10 +58,18 @@ def __post_init__(self) -> None: if not getattr(self, "id", None): self.id = str(uuid.uuid4()) + # Auto-attach the command selected in the UI to user messages the app + # creates during an audio turn (e.g. the transcription in on_audio_end). + # Deserialized payloads are authoritative about their command and reset + # it in from_dict, so a command-less typed message (or resumed history) + # never inherits the audio turn's command. + if self.type == "user_message" and self.command is None: + self.command = context.session.current_command + @classmethod def from_dict(self, _dict: StepDict): type = _dict.get("type", "assistant_message") - return Message( + message = Message( id=_dict["id"], parent_id=_dict.get("parentId"), created_at=_dict["createdAt"], @@ -73,6 +81,14 @@ def from_dict(self, _dict: StepDict): language=_dict.get("language"), metadata=_dict.get("metadata", {}), ) + # A deserialized payload is authoritative about its command: reset it from + # the payload so messages rebuilt here (incoming client messages, thread + # resume) never inherit an active audio turn's command that + # __post_init__ applies to command-less user messages. Normalize like + # __init__ so a non-string payload command can't bypass the str contract. + command = _dict.get("command") + message.command = str(command) if command else None + return message def to_dict(self) -> StepDict: _dict: StepDict = { diff --git a/backend/chainlit/session.py b/backend/chainlit/session.py index fab59b31d7..7876474962 100644 --- a/backend/chainlit/session.py +++ b/backend/chainlit/session.py @@ -124,6 +124,9 @@ class BaseSession: thread_id_to_resume: Optional[str] = None client_type: ClientType current_task: Optional[asyncio.Task] = None + current_command: Optional[str] = None + # Bumped on each audio_start so audio_end won't clear a newer turn's command. + audio_turn: int = 0 chat_started: bool = False def __init__( diff --git a/backend/chainlit/socket.py b/backend/chainlit/socket.py index 74306005ff..f993a0b6e1 100644 --- a/backend/chainlit/socket.py +++ b/backend/chainlit/socket.py @@ -428,16 +428,23 @@ async def window_message(sid, data): @sio.on("audio_start") # pyright: ignore [reportOptionalCall] -async def audio_start(sid): +async def audio_start(sid, payload=None): """Handle audio init.""" session = WebsocketSession.require(sid) context = init_ws_context(session) config: ChainlitConfig = session.get_config() # type: ignore + # Only keep the UI-selected command when audio is accepted (consumed in + # Message.__post_init__), so a declined/disabled start leaves nothing stale. + session.audio_turn += 1 + session.current_command = None + if config.features.audio and config.features.audio.enabled: connected = bool(await config.code.on_audio_start()) connection_state = "on" if connected else "off" + if connected: + session.current_command = payload.get("command") if payload else None await context.emitter.update_audio_connection(connection_state) @@ -462,6 +469,9 @@ async def audio_chunk(sid, payload: InputAudioChunkPayload): async def audio_end(sid): """Handle the end of the audio stream.""" session = WebsocketSession.require(sid) + # Remember which turn we're ending so a dictation started while on_audio_end + # runs keeps its command instead of being cleared out from under it. + audio_turn = session.audio_turn try: context = init_ws_context(session) @@ -484,6 +494,9 @@ async def audio_end(sid): author="Error", content=str(e) or e.__class__.__name__ ).send() finally: + # The command only applies to this audio turn; leave a newer turn's alone. + if session.audio_turn == audio_turn: + session.current_command = None await context.emitter.task_end() diff --git a/backend/tests/test_message.py b/backend/tests/test_message.py index 952f557414..6423c84703 100644 --- a/backend/tests/test_message.py +++ b/backend/tests/test_message.py @@ -24,6 +24,9 @@ def mock_chainlit_context(session=None): mock_loop = Mock(spec=asyncio.AbstractEventLoop) mock_session = session or Mock() mock_session.thread_id = "thread_123" + if session is None: + # Mirror a real session where no command is selected by default. + mock_session.current_command = None with patch("asyncio.get_running_loop", return_value=mock_loop): mock_emitter = AsyncMock() @@ -755,3 +758,90 @@ def test_message_to_dict_with_none_metadata(self): result = msg.to_dict() assert result["metadata"] == {} + + +class TestUserMessageCommandAutoAttach: + """Auto-attaching the session command to user messages (e.g. from audio).""" + + @staticmethod + def _session_with_command(command): + session = Mock() + session.thread_id = "thread_123" + session.current_command = command + return session + + def test_user_message_auto_attaches_current_command(self): + """A user message without a command inherits the session command.""" + session = self._session_with_command("search") + with mock_chainlit_context(session=session): + msg = Message(content="hello", type="user_message") + assert msg.command == "search" + + def test_user_message_keeps_explicit_command(self): + """An explicitly provided command is never overridden.""" + session = self._session_with_command("search") + with mock_chainlit_context(session=session): + msg = Message(content="hello", type="user_message", command="picture") + assert msg.command == "picture" + + def test_user_message_without_current_command_stays_none(self): + """No session command means the user message command stays None.""" + session = self._session_with_command(None) + with mock_chainlit_context(session=session): + msg = Message(content="hello", type="user_message") + assert msg.command is None + + def test_assistant_message_does_not_auto_attach_command(self): + """Only user messages inherit the session command.""" + session = self._session_with_command("search") + with mock_chainlit_context(session=session): + msg = Message(content="hello", type="assistant_message") + assert msg.command is None + + def test_from_dict_user_message_does_not_inherit_session_command(self): + """Deserialized payloads (client messages, resume) never inherit it. + + Guards against a concurrent command-less typed message (or resumed + history) picking up the active audio turn's command. + """ + session = self._session_with_command("search") + step_dict = { + "id": "00000000-0000-4000-8000-000000000000", + "createdAt": "2024-01-01T00:00:00Z", + "output": "typed while an audio turn was active", + "name": "User", + "type": "user_message", + } + with mock_chainlit_context(session=session): + msg = MessageBase.from_dict(step_dict) + assert msg.command is None + + def test_from_dict_user_message_keeps_payload_command(self): + """A deserialized command stays authoritative over the session command.""" + session = self._session_with_command("search") + step_dict = { + "id": "00000000-0000-4000-8000-000000000000", + "createdAt": "2024-01-01T00:00:00Z", + "output": "typed with a different command", + "name": "User", + "type": "user_message", + "command": "picture", + } + with mock_chainlit_context(session=session): + msg = MessageBase.from_dict(step_dict) + assert msg.command == "picture" + + def test_from_dict_normalizes_non_string_command(self): + """A non-string payload command is coerced to str, matching __init__.""" + session = self._session_with_command(None) + step_dict = { + "id": "00000000-0000-4000-8000-000000000000", + "createdAt": "2024-01-01T00:00:00Z", + "output": "command sent as a non-string", + "name": "User", + "type": "user_message", + "command": {"unexpected": "object"}, + } + with mock_chainlit_context(session=session): + msg = MessageBase.from_dict(step_dict) + assert msg.command == str({"unexpected": "object"}) diff --git a/backend/tests/test_socket.py b/backend/tests/test_socket.py index e20959c6d6..d908446cfb 100644 --- a/backend/tests/test_socket.py +++ b/backend/tests/test_socket.py @@ -8,6 +8,8 @@ _authenticate_connection, _get_token, _get_token_from_cookie, + audio_end, + audio_start, clean_session, connection_successful, load_user_env, @@ -628,3 +630,100 @@ async def test_on_chat_start_not_duplicated_on_fresh_then_reconnect( await connection_successful("sid-1") assert on_chat_start.call_count == 1 + + +class TestAudioStartCommand: + """audio_start only keeps the UI-selected command when audio is accepted.""" + + async def _run(self, *, accepted, enabled=True, payload=None): + session = Mock() + session.current_command = "stale" + session.audio_turn = 0 + config = Mock() + config.features.audio.enabled = enabled + config.code.on_audio_start = AsyncMock(return_value=accepted) + session.get_config.return_value = config + + context = Mock() + context.emitter.update_audio_connection = AsyncMock() + + with ( + patch("chainlit.socket.WebsocketSession") as mock_ws, + patch("chainlit.socket.init_ws_context", return_value=context), + ): + mock_ws.require.return_value = session + await audio_start("sid", payload) + return session + + @pytest.mark.asyncio + async def test_accepted_start_keeps_command(self): + session = await self._run(accepted=True, payload={"command": "search"}) + assert session.current_command == "search" + + @pytest.mark.asyncio + async def test_declined_start_clears_command(self): + session = await self._run(accepted=False, payload={"command": "search"}) + assert session.current_command is None + + @pytest.mark.asyncio + async def test_disabled_audio_clears_command(self): + session = await self._run( + accepted=True, enabled=False, payload={"command": "search"} + ) + assert session.current_command is None + + @pytest.mark.asyncio + async def test_accepted_start_without_payload_clears_command(self): + # An older frontend may send audio_start with no payload at all. + session = await self._run(accepted=True, payload=None) + assert session.current_command is None + + @pytest.mark.asyncio + async def test_accepted_start_empty_payload_clears_command(self): + session = await self._run(accepted=True, payload={}) + assert session.current_command is None + + +class TestAudioTurnGuard: + """audio_end clears the command only for the turn it actually ended.""" + + @pytest.mark.asyncio + async def test_older_audio_end_keeps_newer_turn_command(self): + """A slow on_audio_end must not wipe a command a newer turn already set.""" + session = Mock() + session.audio_turn = 0 + session.current_command = None + session.has_first_interaction = True + + config = Mock() + config.features.audio.enabled = True + config.code.on_audio_start = AsyncMock(return_value=True) + session.get_config.return_value = config + + context = Mock() + context.emitter.task_start = AsyncMock() + context.emitter.task_end = AsyncMock() + context.emitter.update_audio_connection = AsyncMock() + + with ( + patch("chainlit.socket.WebsocketSession") as mock_ws, + patch("chainlit.socket.init_ws_context", return_value=context), + ): + mock_ws.require.return_value = session + + # Turn A owns the command. + await audio_start("sid", {"command": "A"}) + assert session.current_command == "A" + + # Turn B starts while turn A's on_audio_end is still running. + async def start_newer_turn(): + await audio_start("sid", {"command": "B"}) + + config.code.on_audio_end = AsyncMock(side_effect=start_newer_turn) + await audio_end("sid") + assert session.current_command == "B" # older audio_end left it alone + + # Turn B's own audio_end owns the turn and clears it. + config.code.on_audio_end = AsyncMock() + await audio_end("sid") + assert session.current_command is None diff --git a/frontend/src/components/chat/MessageComposer/VoiceButton.tsx b/frontend/src/components/chat/MessageComposer/VoiceButton.tsx index 9d4c4ac1bc..345ea0b144 100644 --- a/frontend/src/components/chat/MessageComposer/VoiceButton.tsx +++ b/frontend/src/components/chat/MessageComposer/VoiceButton.tsx @@ -1,5 +1,6 @@ import { X } from 'lucide-react'; import { useHotkeys } from 'react-hotkeys-hook'; +import { useRecoilValue } from 'recoil'; import { useAudio, useConfig } from '@chainlit/react-client'; @@ -12,6 +13,8 @@ import { } from '@/components/ui/tooltip'; import { Translator } from 'components/i18n'; +import { persistentCommandState } from '@/state/chat'; + import { Loader } from '../../Loader'; import { VoiceLines } from '../../icons/VoiceLines'; import { Button } from '../../ui/button'; @@ -23,6 +26,7 @@ interface Props { const VoiceButton = ({ disabled }: Props) => { const { config } = useConfig(); const { startConversation, endConversation, audioConnection } = useAudio(); + const selectedCommand = useRecoilValue(persistentCommandState); const isEnabled = !!config?.features.audio.enabled; useHotkeys( @@ -56,13 +60,19 @@ const VoiceButton = ({ disabled }: Props) => { } if (audioConnection === 'on') return endConversation(); - return startConversation(); + return startConversation(selectedCommand?.id); }, { enableOnFormTags: false, preventDefault: false // Don't prevent default - let letters be typed }, - [isEnabled, audioConnection, startConversation, endConversation] + [ + isEnabled, + audioConnection, + startConversation, + endConversation, + selectedCommand + ] ); if (!isEnabled) return null; @@ -90,7 +100,7 @@ const VoiceButton = ({ disabled }: Props) => { audioConnection === 'on' ? endConversation : audioConnection === 'off' - ? startConversation + ? () => startConversation(selectedCommand?.id) : undefined } > diff --git a/libs/react-client/src/useAudio.ts b/libs/react-client/src/useAudio.ts index 79a964f3f2..6bcdef3c25 100644 --- a/libs/react-client/src/useAudio.ts +++ b/libs/react-client/src/useAudio.ts @@ -18,10 +18,13 @@ const useAudio = () => { const { startAudioStream, endAudioStream } = useChatInteract(); - const startConversation = useCallback(async () => { - setAudioConnection('connecting'); - await startAudioStream(); - }, [startAudioStream]); + const startConversation = useCallback( + async (command?: string) => { + setAudioConnection('connecting'); + await startAudioStream(command); + }, + [startAudioStream] + ); const endConversation = useCallback(async () => { setAudioConnection('off'); diff --git a/libs/react-client/src/useChatInteract.ts b/libs/react-client/src/useChatInteract.ts index 7436d7786b..93c472e5e9 100644 --- a/libs/react-client/src/useChatInteract.ts +++ b/libs/react-client/src/useChatInteract.ts @@ -133,9 +133,12 @@ const useChatInteract = () => { [session?.socket] ); - const startAudioStream = useCallback(() => { - session?.socket.emit('audio_start'); - }, [session?.socket]); + const startAudioStream = useCallback( + (command?: string) => { + session?.socket.emit('audio_start', { command }); + }, + [session?.socket] + ); const sendAudioChunk = useCallback( (