Sync from GitHub via hub-sync
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .env.example +19 -6
- README.md +58 -23
- profiles/bored_teenager/tools.txt +1 -1
- profiles/captain_circuit/tools.txt +1 -1
- profiles/chess_coach/tools.txt +1 -1
- profiles/cosmic_kitchen/tools.txt +1 -1
- profiles/default/tools.txt +1 -1
- profiles/example/tools.txt +1 -1
- profiles/hype_bot/tools.txt +1 -1
- profiles/mad_scientist_assistant/tools.txt +1 -1
- profiles/mars_rover/tools.txt +1 -1
- profiles/nature_documentarian/tools.txt +1 -1
- profiles/noir_detective/tools.txt +1 -1
- profiles/sorry_bro/tools.txt +1 -1
- profiles/tedai/tools.txt +1 -1
- profiles/time_traveler/tools.txt +1 -1
- profiles/victorian_butler/tools.txt +1 -1
- pyproject.toml +3 -2
- src/reachy_mini_conversation_app/audio/startup_config.py +64 -0
- src/reachy_mini_conversation_app/base_realtime.py +1017 -0
- src/reachy_mini_conversation_app/config.py +205 -5
- src/reachy_mini_conversation_app/console.py +163 -72
- src/reachy_mini_conversation_app/conversation_handler.py +69 -0
- src/reachy_mini_conversation_app/gemini_live.py +40 -9
- src/reachy_mini_conversation_app/gradio_personality.py +56 -37
- src/reachy_mini_conversation_app/headless_personality.py +14 -3
- src/reachy_mini_conversation_app/headless_personality_ui.py +12 -16
- src/reachy_mini_conversation_app/huggingface_realtime.py +160 -0
- src/reachy_mini_conversation_app/main.py +89 -19
- src/reachy_mini_conversation_app/openai_realtime.py +136 -876
- src/reachy_mini_conversation_app/startup_settings.py +106 -0
- src/reachy_mini_conversation_app/static/index.html +53 -9
- src/reachy_mini_conversation_app/static/main.js +206 -21
- src/reachy_mini_conversation_app/static/style.css +24 -1
- src/reachy_mini_conversation_app/tools/core_tools.py +8 -8
- src/reachy_mini_conversation_app/tools/dance.py +25 -25
- src/reachy_mini_conversation_app/tools/{do_nothing.py → idle_do_nothing.py} +12 -9
- src/reachy_mini_conversation_app/tools/play_emotion.py +14 -5
- src/reachy_mini_conversation_app/vision/head_tracking/yolo.py +12 -2
- src/reachy_mini_conversation_app/vision/head_tracking/yolo_process.py +10 -2
- src/reachy_mini_conversation_app/vision/local_vision.py +72 -17
- tests/audio/test_head_wobbler.py +6 -1
- tests/audio/test_startup_config.py +95 -0
- tests/test_config_name_collisions.py +83 -0
- tests/test_console.py +359 -10
- tests/test_gemini_live.py +96 -15
- tests/test_huggingface_realtime.py +628 -0
- tests/test_openai_realtime.py +543 -54
- tests/test_profile_paths.py +42 -0
- tests/test_startup_settings.py +80 -0
.env.example
CHANGED
|
@@ -1,12 +1,24 @@
|
|
| 1 |
-
# Realtime backend to use: "openai" or "gemini"
|
| 2 |
-
|
| 3 |
-
#
|
| 4 |
-
|
|
|
|
|
|
|
| 5 |
|
| 6 |
# Set up your API key according to your backend
|
| 7 |
OPENAI_API_KEY=
|
| 8 |
# GEMINI_API_KEY=
|
| 9 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 10 |
# Local vision model (only used with --local-vision CLI flag)
|
| 11 |
# By default, vision is handled by the selected realtime backend when the camera tool is used
|
| 12 |
LOCAL_VISION_MODEL=HuggingFaceTB/SmolVLM2-2.2B-Instruct
|
|
@@ -17,8 +29,9 @@ HF_HOME=./cache
|
|
| 17 |
# Hugging Face token for accessing datasets/models
|
| 18 |
HF_TOKEN=
|
| 19 |
|
| 20 |
-
#
|
| 21 |
-
|
|
|
|
| 22 |
|
| 23 |
# Optional external profile/tool directories
|
| 24 |
# REACHY_MINI_EXTERNAL_PROFILES_DIRECTORY=external_content/external_profiles
|
|
|
|
| 1 |
+
# Realtime backend to use: "huggingface", "openai", or "gemini"
|
| 2 |
+
# Defaults to "huggingface" when unset.
|
| 3 |
+
# BACKEND_PROVIDER="huggingface"
|
| 4 |
+
# Optional model override for OpenAI Realtime or Gemini Live.
|
| 5 |
+
# Hugging Face uses the server's model selection and ignores MODEL_NAME.
|
| 6 |
+
# MODEL_NAME="gpt-realtime"
|
| 7 |
|
| 8 |
# Set up your API key according to your backend
|
| 9 |
OPENAI_API_KEY=
|
| 10 |
# GEMINI_API_KEY=
|
| 11 |
|
| 12 |
+
# Hugging Face connection mode: "deployed" or "local".
|
| 13 |
+
# Deployed mode uses the built-in Hugging Face server.
|
| 14 |
+
# Defaults to "deployed" when unset.
|
| 15 |
+
# HF_REALTIME_CONNECTION_MODE="deployed"
|
| 16 |
+
|
| 17 |
+
# Direct Hugging Face realtime endpoint for local/LAN backends.
|
| 18 |
+
# Accepts either the base URL (`ws://127.0.0.1:8765/v1`) or the full realtime URL
|
| 19 |
+
# (`ws://127.0.0.1:8765/v1/realtime`). Used when HF_REALTIME_CONNECTION_MODE=local.
|
| 20 |
+
# HF_REALTIME_WS_URL="ws://127.0.0.1:8765/v1/realtime"
|
| 21 |
+
|
| 22 |
# Local vision model (only used with --local-vision CLI flag)
|
| 23 |
# By default, vision is handled by the selected realtime backend when the camera tool is used
|
| 24 |
LOCAL_VISION_MODEL=HuggingFaceTB/SmolVLM2-2.2B-Instruct
|
|
|
|
| 29 |
# Hugging Face token for accessing datasets/models
|
| 30 |
HF_TOKEN=
|
| 31 |
|
| 32 |
+
# Optional built-in profile selector. Use a folder name from profiles/.
|
| 33 |
+
# If unset, the app uses the default profile unless startup settings were saved from the UI.
|
| 34 |
+
# REACHY_MINI_CUSTOM_PROFILE=mars_rover
|
| 35 |
|
| 36 |
# Optional external profile/tool directories
|
| 37 |
# REACHY_MINI_EXTERNAL_PROFILES_DIRECTORY=external_content/external_profiles
|
README.md
CHANGED
|
@@ -14,7 +14,7 @@ tags:
|
|
| 14 |
|
| 15 |
# Reachy Mini conversation app
|
| 16 |
|
| 17 |
-
Conversational app for the Reachy Mini robot combining
|
| 18 |
|
| 19 |

|
| 20 |
|
|
@@ -30,10 +30,11 @@ Conversational app for the Reachy Mini robot combining real-time voice APIs (Ope
|
|
| 30 |
- [License](#license)
|
| 31 |
|
| 32 |
## Overview
|
| 33 |
-
- Real-time audio conversation loop with `fastrtc` for low-latency streaming.
|
| 34 |
-
- **
|
| 35 |
-
- **
|
| 36 |
-
-
|
|
|
|
| 37 |
- Layered motion system queues primary moves (dances, emotions, goto poses, breathing) while blending speech-reactive wobble and head-tracking.
|
| 38 |
- Async tool dispatch integrates robot motion, camera capture, and optional head-tracking capabilities through a Gradio web UI with live transcripts.
|
| 39 |
|
|
@@ -120,33 +121,64 @@ Some wheels (like PyTorch) are large and require compatible CUDA or CPU builds
|
|
| 120 |
|
| 121 |
## Configuration
|
| 122 |
|
| 123 |
-
|
| 124 |
-
|
|
|
|
| 125 |
|
| 126 |
| Variable | Description |
|
| 127 |
|----------|-------------|
|
| 128 |
-
| `OPENAI_API_KEY` |
|
| 129 |
| `GEMINI_API_KEY` | Required for Gemini mode. Also accepts `GOOGLE_API_KEY`. Get one at [aistudio.google.com](https://aistudio.google.com/apikey). |
|
| 130 |
-
| `BACKEND_PROVIDER` | Realtime backend to use: `
|
| 131 |
-
| `MODEL_NAME` | Optional model override for
|
|
|
|
|
|
|
| 132 |
| `HF_HOME` | Cache directory for local Hugging Face downloads (only used with `--local-vision` flag, defaults to `./cache`). |
|
| 133 |
| `HF_TOKEN` | Optional token for Hugging Face access (for gated/private assets). |
|
| 134 |
| `LOCAL_VISION_MODEL` | Hugging Face model path for local vision processing (only used with `--local-vision` flag, defaults to `HuggingFaceTB/SmolVLM2-2.2B-Instruct`). |
|
| 135 |
|
| 136 |
-
###
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 137 |
|
| 138 |
-
|
| 139 |
|
| 140 |
```env
|
| 141 |
-
BACKEND_PROVIDER=
|
| 142 |
-
|
| 143 |
-
|
| 144 |
```
|
| 145 |
|
| 146 |
-
|
| 147 |
|
| 148 |
-
|
| 149 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 150 |
|
| 151 |
## Running the app
|
| 152 |
|
|
@@ -167,7 +199,7 @@ The app runs in console mode by default. Add `--gradio` to launch a web UI at ht
|
|
| 167 |
|--------|---------|-------------|
|
| 168 |
| `--head-tracker {yolo,mediapipe}` | `None` | Select a head-tracking backend when a camera is available. `yolo` uses a local YOLO face detector, `mediapipe` comes from the `reachy_mini_toolbox` package. Requires the matching optional extra. |
|
| 169 |
| `--no-camera` | `False` | Run without camera capture or head tracking. |
|
| 170 |
-
| `--local-vision` | `False` | Use the local vision model (SmolVLM2) for camera-tool requests instead of the selected realtime backend
|
| 171 |
| `--gradio` | `False` | Launch the Gradio web UI. Without this flag, runs in console mode. Required when running in simulation mode. |
|
| 172 |
| `--robot-name` | `None` | Optional. Connect to a specific robot by name when running multiple daemons on the same subnet. See [Multiple robots on the same subnet](#advanced-features). |
|
| 173 |
| `--debug` | `False` | Enable verbose logging for troubleshooting. |
|
|
@@ -205,7 +237,7 @@ reachy-mini-conversation-app --gradio
|
|
| 205 |
| `stop_dance` | Clear queued dances. | Core install only. |
|
| 206 |
| `play_emotion` | Play a recorded emotion clip via Hugging Face datasets. | Core install only. Uses the default open emotions dataset: [`pollen-robotics/reachy-mini-emotions-library`](https://huggingface.co/datasets/pollen-robotics/reachy-mini-emotions-library). |
|
| 207 |
| `stop_emotion` | Clear queued emotions. | Core install only. |
|
| 208 |
-
| `
|
| 209 |
|
| 210 |
## Advanced features
|
| 211 |
|
|
@@ -218,7 +250,9 @@ Built-in motion content is published as open Hugging Face datasets:
|
|
| 218 |
|
| 219 |
Create custom profiles with dedicated instructions and enabled tools.
|
| 220 |
|
| 221 |
-
|
|
|
|
|
|
|
| 222 |
|
| 223 |
Each profile should include `instructions.txt` (prompt text). `tools.txt` (list of allowed tools) is recommended. If missing for a non-default profile, the app falls back to `profiles/default/tools.txt`. Profiles can optionally contain custom tool implementations.
|
| 224 |
|
|
@@ -266,7 +300,7 @@ To create a locked variant of the app that cannot switch profiles, edit `src/rea
|
|
| 266 |
```python
|
| 267 |
LOCKED_PROFILE: str | None = "mars_rover" # Lock to this profile
|
| 268 |
```
|
| 269 |
-
When `LOCKED_PROFILE` is set, the app always uses that profile, ignoring `REACHY_MINI_CUSTOM_PROFILE`
|
| 270 |
This is useful for creating dedicated clones of the app with a fixed personality. Clone scripts can simply edit this constant to lock the variant.
|
| 271 |
|
| 272 |
</details>
|
|
@@ -294,9 +328,10 @@ external_content/
|
|
| 294 |
|
| 295 |
**Environment variables:**
|
| 296 |
|
| 297 |
-
Set these values in your `.env`
|
| 298 |
|
| 299 |
```env
|
|
|
|
| 300 |
REACHY_MINI_CUSTOM_PROFILE=my_profile
|
| 301 |
REACHY_MINI_EXTERNAL_PROFILES_DIRECTORY=./external_content/external_profiles
|
| 302 |
REACHY_MINI_EXTERNAL_TOOLS_DIRECTORY=./external_content/external_tools
|
|
|
|
| 14 |
|
| 15 |
# Reachy Mini conversation app
|
| 16 |
|
| 17 |
+
Conversational app for the Reachy Mini robot combining realtime voice backends, vision pipelines, and choreographed motion libraries.
|
| 18 |
|
| 19 |

|
| 20 |
|
|
|
|
| 30 |
- [License](#license)
|
| 31 |
|
| 32 |
## Overview
|
| 33 |
+
- Real-time audio conversation loop with `fastrtc` for low-latency streaming. Supported backends:
|
| 34 |
+
- **Hugging Face** - default, using the built-in Hugging Face server or your own local endpoint.
|
| 35 |
+
- **OpenAI Realtime** (`gpt-realtime`) - requires `OPENAI_API_KEY`.
|
| 36 |
+
- **Gemini Live** (`gemini-3.1-flash-live-preview`) - requires `GEMINI_API_KEY`.
|
| 37 |
+
- Vision processing uses the selected realtime backend by default (when the camera tool is used), with optional on-device local vision using SmolVLM2 (CPU/GPU/MPS) via `--local-vision`.
|
| 38 |
- Layered motion system queues primary moves (dances, emotions, goto poses, breathing) while blending speech-reactive wobble and head-tracking.
|
| 39 |
- Async tool dispatch integrates robot motion, camera capture, and optional head-tracking capabilities through a Gradio web UI with live transcripts.
|
| 40 |
|
|
|
|
| 121 |
|
| 122 |
## Configuration
|
| 123 |
|
| 124 |
+
The default setup uses the Hugging Face backend and does not require an API key.
|
| 125 |
+
|
| 126 |
+
Copy `.env.example` to `.env` when you want to switch backends, provide API keys, or point Hugging Face at your own local endpoint.
|
| 127 |
|
| 128 |
| Variable | Description |
|
| 129 |
|----------|-------------|
|
| 130 |
+
| `OPENAI_API_KEY` | Required for OpenAI Realtime mode. |
|
| 131 |
| `GEMINI_API_KEY` | Required for Gemini mode. Also accepts `GOOGLE_API_KEY`. Get one at [aistudio.google.com](https://aistudio.google.com/apikey). |
|
| 132 |
+
| `BACKEND_PROVIDER` | Realtime backend to use: `huggingface` (default), `openai`, or `gemini`. |
|
| 133 |
+
| `MODEL_NAME` | Optional model override for OpenAI Realtime or Gemini Live. Defaults to `gpt-realtime` for OpenAI and `gemini-3.1-flash-live-preview` for Gemini. Hugging Face uses the server's model selection. |
|
| 134 |
+
| `HF_REALTIME_CONNECTION_MODE` | Hugging Face connection selector: `deployed` uses the built-in Hugging Face server; `local` uses `HF_REALTIME_WS_URL`. Defaults to `deployed`. |
|
| 135 |
+
| `HF_REALTIME_WS_URL` | Direct websocket endpoint for your own Hugging Face backend. Accepts either a base URL like `ws://127.0.0.1:8765/v1` or the full websocket URL `ws://127.0.0.1:8765/v1/realtime`. Used when `HF_REALTIME_CONNECTION_MODE=local`. |
|
| 136 |
| `HF_HOME` | Cache directory for local Hugging Face downloads (only used with `--local-vision` flag, defaults to `./cache`). |
|
| 137 |
| `HF_TOKEN` | Optional token for Hugging Face access (for gated/private assets). |
|
| 138 |
| `LOCAL_VISION_MODEL` | Hugging Face model path for local vision processing (only used with `--local-vision` flag, defaults to `HuggingFaceTB/SmolVLM2-2.2B-Instruct`). |
|
| 139 |
|
| 140 |
+
### Hugging Face Connection Modes
|
| 141 |
+
|
| 142 |
+
Use the built-in Hugging Face server through the app-managed Space proxy. This is the default for a new install; set it explicitly only when you want to switch back from a saved local endpoint:
|
| 143 |
+
|
| 144 |
+
```env
|
| 145 |
+
BACKEND_PROVIDER=huggingface
|
| 146 |
+
HF_REALTIME_CONNECTION_MODE=deployed
|
| 147 |
+
```
|
| 148 |
+
|
| 149 |
+
Run your own realtime voice backend using [speech-to-speech](https://github.com/huggingface/speech-to-speech) on the same machine as the conversation app:
|
| 150 |
+
|
| 151 |
+
```env
|
| 152 |
+
BACKEND_PROVIDER=huggingface
|
| 153 |
+
HF_REALTIME_CONNECTION_MODE=local
|
| 154 |
+
HF_REALTIME_WS_URL=ws://127.0.0.1:8765/v1/realtime
|
| 155 |
+
```
|
| 156 |
|
| 157 |
+
Run your own Hugging Face backend on your laptop and connect to it from Reachy Mini Wireless over the same Wi-Fi network:
|
| 158 |
|
| 159 |
```env
|
| 160 |
+
BACKEND_PROVIDER=huggingface
|
| 161 |
+
HF_REALTIME_CONNECTION_MODE=local
|
| 162 |
+
HF_REALTIME_WS_URL=ws://<your-laptop-lan-ip>:8765/v1/realtime
|
| 163 |
```
|
| 164 |
|
| 165 |
+
For that LAN setup, make sure the backend listens on an address reachable from the robot, not only on `127.0.0.1`.
|
| 166 |
|
| 167 |
+
If the backend stays bound to loopback on your laptop, you can forward it into the robot over SSH instead:
|
| 168 |
+
|
| 169 |
+
```bash
|
| 170 |
+
ssh -N -R 8765:127.0.0.1:8765 <robot-user>@<robot-host>
|
| 171 |
+
```
|
| 172 |
+
|
| 173 |
+
Then set this on the robot:
|
| 174 |
+
|
| 175 |
+
```env
|
| 176 |
+
BACKEND_PROVIDER=huggingface
|
| 177 |
+
HF_REALTIME_CONNECTION_MODE=local
|
| 178 |
+
HF_REALTIME_WS_URL=ws://127.0.0.1:8765/v1/realtime
|
| 179 |
+
```
|
| 180 |
+
|
| 181 |
+
When using the headless settings UI, selecting `Hugging Face` lets you choose either the built-in server or a local `host:port` target. The UI writes `HF_REALTIME_CONNECTION_MODE` for you, and the local path writes `HF_REALTIME_WS_URL` with a default of `localhost:8765`.
|
| 182 |
|
| 183 |
## Running the app
|
| 184 |
|
|
|
|
| 199 |
|--------|---------|-------------|
|
| 200 |
| `--head-tracker {yolo,mediapipe}` | `None` | Select a head-tracking backend when a camera is available. `yolo` uses a local YOLO face detector, `mediapipe` comes from the `reachy_mini_toolbox` package. Requires the matching optional extra. |
|
| 201 |
| `--no-camera` | `False` | Run without camera capture or head tracking. |
|
| 202 |
+
| `--local-vision` | `False` | Use the local vision model (SmolVLM2) for camera-tool requests instead of the selected realtime backend. Requires `local_vision` extra to be installed. |
|
| 203 |
| `--gradio` | `False` | Launch the Gradio web UI. Without this flag, runs in console mode. Required when running in simulation mode. |
|
| 204 |
| `--robot-name` | `None` | Optional. Connect to a specific robot by name when running multiple daemons on the same subnet. See [Multiple robots on the same subnet](#advanced-features). |
|
| 205 |
| `--debug` | `False` | Enable verbose logging for troubleshooting. |
|
|
|
|
| 237 |
| `stop_dance` | Clear queued dances. | Core install only. |
|
| 238 |
| `play_emotion` | Play a recorded emotion clip via Hugging Face datasets. | Core install only. Uses the default open emotions dataset: [`pollen-robotics/reachy-mini-emotions-library`](https://huggingface.co/datasets/pollen-robotics/reachy-mini-emotions-library). |
|
| 239 |
| `stop_emotion` | Clear queued emotions. | Core install only. |
|
| 240 |
+
| `idle_do_nothing` | Explicitly remain idle during an idle turn. Not intended for normal conversation turns. | Core install only. |
|
| 241 |
|
| 242 |
## Advanced features
|
| 243 |
|
|
|
|
| 250 |
|
| 251 |
Create custom profiles with dedicated instructions and enabled tools.
|
| 252 |
|
| 253 |
+
For normal usage, select a profile from the UI and save it for startup. That selection is persisted in `startup_settings.json`.
|
| 254 |
+
|
| 255 |
+
If no startup settings have been saved yet, you can still seed startup from the environment with `REACHY_MINI_CUSTOM_PROFILE=<name>` to load `profiles/<name>/`. If neither is set, the `default` profile is used.
|
| 256 |
|
| 257 |
Each profile should include `instructions.txt` (prompt text). `tools.txt` (list of allowed tools) is recommended. If missing for a non-default profile, the app falls back to `profiles/default/tools.txt`. Profiles can optionally contain custom tool implementations.
|
| 258 |
|
|
|
|
| 300 |
```python
|
| 301 |
LOCKED_PROFILE: str | None = "mars_rover" # Lock to this profile
|
| 302 |
```
|
| 303 |
+
When `LOCKED_PROFILE` is set, the app always uses that profile, ignoring saved startup settings, `REACHY_MINI_CUSTOM_PROFILE`, and the Gradio UI. The UI shows "(locked)" and disables all profile editing controls.
|
| 304 |
This is useful for creating dedicated clones of the app with a fixed personality. Clone scripts can simply edit this constant to lock the variant.
|
| 305 |
|
| 306 |
</details>
|
|
|
|
| 328 |
|
| 329 |
**Environment variables:**
|
| 330 |
|
| 331 |
+
Set these values in your `.env` when you want env-driven external profile/tool selection:
|
| 332 |
|
| 333 |
```env
|
| 334 |
+
# Optional fallback/manual profile selector:
|
| 335 |
REACHY_MINI_CUSTOM_PROFILE=my_profile
|
| 336 |
REACHY_MINI_EXTERNAL_PROFILES_DIRECTORY=./external_content/external_profiles
|
| 337 |
REACHY_MINI_EXTERNAL_TOOLS_DIRECTORY=./external_content/external_tools
|
profiles/bored_teenager/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/captain_circuit/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/chess_coach/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/cosmic_kitchen/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/default/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/example/tools.txt
CHANGED
|
@@ -5,7 +5,7 @@ stop_dance
|
|
| 5 |
play_emotion
|
| 6 |
stop_emotion
|
| 7 |
# camera
|
| 8 |
-
#
|
| 9 |
# head_tracking
|
| 10 |
# move_head
|
| 11 |
|
|
|
|
| 5 |
play_emotion
|
| 6 |
stop_emotion
|
| 7 |
# camera
|
| 8 |
+
# idle_do_nothing
|
| 9 |
# head_tracking
|
| 10 |
# move_head
|
| 11 |
|
profiles/hype_bot/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/mad_scientist_assistant/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/mars_rover/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/nature_documentarian/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/noir_detective/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/sorry_bro/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/tedai/tools.txt
CHANGED
|
@@ -3,5 +3,5 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
move_head
|
profiles/time_traveler/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
profiles/victorian_butler/tools.txt
CHANGED
|
@@ -3,6 +3,6 @@ stop_dance
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
-
|
| 7 |
head_tracking
|
| 8 |
move_head
|
|
|
|
| 3 |
play_emotion
|
| 4 |
stop_emotion
|
| 5 |
camera
|
| 6 |
+
idle_do_nothing
|
| 7 |
head_tracking
|
| 8 |
move_head
|
pyproject.toml
CHANGED
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
| 4 |
|
| 5 |
[project]
|
| 6 |
name = "reachy_mini_conversation_app"
|
| 7 |
-
version = "0.
|
| 8 |
authors = [{ name = "Pollen Robotics", email = "contact@pollen-robotics.com" }]
|
| 9 |
description = ""
|
| 10 |
readme = "README.md"
|
|
@@ -16,6 +16,7 @@ dependencies = [
|
|
| 16 |
"fastrtc>=0.0.34",
|
| 17 |
"gradio==5.50.1.dev1",
|
| 18 |
"huggingface-hub==1.3.0",
|
|
|
|
| 19 |
|
| 20 |
#Environment variables
|
| 21 |
"python-dotenv",
|
|
@@ -29,7 +30,7 @@ dependencies = [
|
|
| 29 |
#Reachy mini
|
| 30 |
"reachy_mini_dances_library",
|
| 31 |
"reachy_mini_toolbox",
|
| 32 |
-
"reachy-mini>=1.
|
| 33 |
"gradio_client>=1.13.3",
|
| 34 |
]
|
| 35 |
|
|
|
|
| 4 |
|
| 5 |
[project]
|
| 6 |
name = "reachy_mini_conversation_app"
|
| 7 |
+
version = "0.6.0"
|
| 8 |
authors = [{ name = "Pollen Robotics", email = "contact@pollen-robotics.com" }]
|
| 9 |
description = ""
|
| 10 |
readme = "README.md"
|
|
|
|
| 16 |
"fastrtc>=0.0.34",
|
| 17 |
"gradio==5.50.1.dev1",
|
| 18 |
"huggingface-hub==1.3.0",
|
| 19 |
+
"httpx",
|
| 20 |
|
| 21 |
#Environment variables
|
| 22 |
"python-dotenv",
|
|
|
|
| 30 |
#Reachy mini
|
| 31 |
"reachy_mini_dances_library",
|
| 32 |
"reachy_mini_toolbox",
|
| 33 |
+
"reachy-mini>=1.7.1",
|
| 34 |
"gradio_client>=1.13.3",
|
| 35 |
]
|
| 36 |
|
src/reachy_mini_conversation_app/audio/startup_config.py
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Startup configuration for the Reachy Mini audio processor."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
import logging
|
| 5 |
+
from collections.abc import Sequence
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
AudioControlValue = float | int
|
| 9 |
+
AudioStartupParameter = tuple[str, tuple[AudioControlValue, ...]]
|
| 10 |
+
WRITE_SETTLE_SECONDS = 0.1
|
| 11 |
+
|
| 12 |
+
AUDIO_STARTUP_CONFIG: tuple[AudioStartupParameter, ...] = (
|
| 13 |
+
("PP_AGCMAXGAIN", (10.0,)),
|
| 14 |
+
("PP_MIN_NS", (0.8,)),
|
| 15 |
+
("PP_MIN_NN", (0.8,)),
|
| 16 |
+
("PP_GAMMA_E", (0.5,)),
|
| 17 |
+
("PP_GAMMA_ETAIL", (0.5,)),
|
| 18 |
+
("PP_NLATTENONOFF", (0,)),
|
| 19 |
+
("PP_MGSCALE", (4.0, 1.0, 1.0)),
|
| 20 |
+
)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def apply_audio_startup_config(
|
| 24 |
+
robot: object,
|
| 25 |
+
*,
|
| 26 |
+
logger: logging.Logger | None = None,
|
| 27 |
+
verify: bool = True,
|
| 28 |
+
write_settle_seconds: float = WRITE_SETTLE_SECONDS,
|
| 29 |
+
) -> bool:
|
| 30 |
+
"""Apply the tuned XVF3800 audio configuration for the conversation app."""
|
| 31 |
+
log = logger or logging.getLogger(__name__)
|
| 32 |
+
audio = getattr(getattr(robot, "media", None), "audio", None)
|
| 33 |
+
|
| 34 |
+
if audio is None:
|
| 35 |
+
log.warning("Skipping Reachy audio startup config: robot media audio is unavailable.")
|
| 36 |
+
return False
|
| 37 |
+
|
| 38 |
+
apply_audio_config = getattr(audio, "apply_audio_config", None)
|
| 39 |
+
if not callable(apply_audio_config):
|
| 40 |
+
log.warning("Skipping Reachy audio startup config: SDK audio config API is unavailable.")
|
| 41 |
+
return False
|
| 42 |
+
|
| 43 |
+
try:
|
| 44 |
+
applied = bool(
|
| 45 |
+
apply_audio_config(
|
| 46 |
+
AUDIO_STARTUP_CONFIG,
|
| 47 |
+
verify=verify,
|
| 48 |
+
write_settle_seconds=write_settle_seconds,
|
| 49 |
+
)
|
| 50 |
+
)
|
| 51 |
+
except Exception as exc:
|
| 52 |
+
log.warning("Skipping Reachy audio startup config: SDK audio config failed: %s", exc)
|
| 53 |
+
return False
|
| 54 |
+
|
| 55 |
+
if applied:
|
| 56 |
+
log.info("Applied Reachy audio startup config: %s", _format_config(AUDIO_STARTUP_CONFIG))
|
| 57 |
+
else:
|
| 58 |
+
log.warning("Reachy audio startup config was not applied.")
|
| 59 |
+
|
| 60 |
+
return applied
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
def _format_config(config: Sequence[AudioStartupParameter]) -> str:
|
| 64 |
+
return ", ".join(f"{name}={' '.join(str(value) for value in values)}" for name, values in config)
|
src/reachy_mini_conversation_app/base_realtime.py
ADDED
|
@@ -0,0 +1,1017 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
import time
|
| 3 |
+
import uuid
|
| 4 |
+
import base64
|
| 5 |
+
import random
|
| 6 |
+
import asyncio
|
| 7 |
+
import logging
|
| 8 |
+
from abc import ABC, abstractmethod
|
| 9 |
+
from typing import Any, Final, Tuple, ClassVar, Optional
|
| 10 |
+
from datetime import datetime
|
| 11 |
+
|
| 12 |
+
import numpy as np
|
| 13 |
+
import gradio as gr
|
| 14 |
+
from openai import AsyncOpenAI
|
| 15 |
+
from fastrtc import AdditionalOutputs, wait_for_item, audio_to_int16
|
| 16 |
+
from pydantic import Field, BaseModel
|
| 17 |
+
from numpy.typing import NDArray
|
| 18 |
+
from scipy.signal import resample
|
| 19 |
+
from openai.types.realtime import (
|
| 20 |
+
RealtimeAudioConfigParam,
|
| 21 |
+
RealtimeToolsConfigParam,
|
| 22 |
+
RealtimeFunctionToolParam,
|
| 23 |
+
RealtimeAudioConfigOutputParam,
|
| 24 |
+
RealtimeResponseCreateParamsParam,
|
| 25 |
+
RealtimeSessionCreateRequestParam,
|
| 26 |
+
)
|
| 27 |
+
from websockets.exceptions import ConnectionClosedError
|
| 28 |
+
from openai.resources.realtime.realtime import AsyncRealtimeConnection
|
| 29 |
+
|
| 30 |
+
from reachy_mini_conversation_app.config import (
|
| 31 |
+
config,
|
| 32 |
+
get_default_voice_for_backend,
|
| 33 |
+
get_available_voices_for_backend,
|
| 34 |
+
)
|
| 35 |
+
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 36 |
+
from reachy_mini_conversation_app.conversation_handler import ConversationHandler
|
| 37 |
+
from reachy_mini_conversation_app.tools.background_tool_manager import (
|
| 38 |
+
ToolCallRoutine,
|
| 39 |
+
ToolNotification,
|
| 40 |
+
BackgroundToolManager,
|
| 41 |
+
)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
logger = logging.getLogger(__name__)
|
| 45 |
+
|
| 46 |
+
_RESPONSE_DONE_TIMEOUT: Final[float] = 30.0
|
| 47 |
+
_RESPONSE_REJECTION_RETRY_DELAY: Final[float] = 0.5
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
class InputTranscriptChunksByItem(BaseModel):
|
| 51 |
+
"""Current item_id and its accumulated deltas. Only one item at a time."""
|
| 52 |
+
|
| 53 |
+
item_id: str | None = None
|
| 54 |
+
deltas: list[str] = Field(default_factory=list)
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def to_realtime_tools_config(tool_specs: list[dict[str, Any]]) -> RealtimeToolsConfigParam:
|
| 58 |
+
"""Convert app tool specs to the OpenAI-compatible realtime session shape."""
|
| 59 |
+
realtime_tools: RealtimeToolsConfigParam = []
|
| 60 |
+
for spec in tool_specs:
|
| 61 |
+
tool_type = spec.get("type")
|
| 62 |
+
name = spec.get("name")
|
| 63 |
+
description = spec.get("description")
|
| 64 |
+
parameters = spec.get("parameters", {})
|
| 65 |
+
|
| 66 |
+
if tool_type != "function" or not isinstance(name, str):
|
| 67 |
+
raise ValueError(f"Unsupported realtime tool spec: {spec!r}")
|
| 68 |
+
|
| 69 |
+
realtime_tool = RealtimeFunctionToolParam(
|
| 70 |
+
type="function",
|
| 71 |
+
name=name,
|
| 72 |
+
parameters=parameters,
|
| 73 |
+
)
|
| 74 |
+
if isinstance(description, str):
|
| 75 |
+
realtime_tool["description"] = description
|
| 76 |
+
realtime_tools.append(realtime_tool)
|
| 77 |
+
return realtime_tools
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
class BaseRealtimeHandler(ConversationHandler, ABC):
|
| 81 |
+
"""Shared realtime stream handler for OpenAI-compatible client APIs."""
|
| 82 |
+
|
| 83 |
+
BACKEND_PROVIDER: ClassVar[str]
|
| 84 |
+
SAMPLE_RATE: ClassVar[int]
|
| 85 |
+
REFRESH_CLIENT_ON_RECONNECT: ClassVar[bool]
|
| 86 |
+
AUDIO_INPUT_COST_PER_1M: ClassVar[float]
|
| 87 |
+
AUDIO_OUTPUT_COST_PER_1M: ClassVar[float]
|
| 88 |
+
TEXT_INPUT_COST_PER_1M: ClassVar[float]
|
| 89 |
+
TEXT_OUTPUT_COST_PER_1M: ClassVar[float]
|
| 90 |
+
IMAGE_INPUT_COST_PER_1M: ClassVar[float]
|
| 91 |
+
|
| 92 |
+
_REQUIRED_PROVIDER_CONFIG: ClassVar[tuple[str, ...]] = (
|
| 93 |
+
"BACKEND_PROVIDER",
|
| 94 |
+
"SAMPLE_RATE",
|
| 95 |
+
"REFRESH_CLIENT_ON_RECONNECT",
|
| 96 |
+
"AUDIO_INPUT_COST_PER_1M",
|
| 97 |
+
"AUDIO_OUTPUT_COST_PER_1M",
|
| 98 |
+
"TEXT_INPUT_COST_PER_1M",
|
| 99 |
+
"TEXT_OUTPUT_COST_PER_1M",
|
| 100 |
+
"IMAGE_INPUT_COST_PER_1M",
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
def __init_subclass__(cls, **kwargs: Any) -> None:
|
| 104 |
+
"""Require concrete providers to declare their provider configuration."""
|
| 105 |
+
super().__init_subclass__(**kwargs)
|
| 106 |
+
missing = [name for name in cls._REQUIRED_PROVIDER_CONFIG if name not in cls.__dict__]
|
| 107 |
+
if missing:
|
| 108 |
+
raise TypeError(f"{cls.__name__} must define provider config class variable(s): {', '.join(missing)}")
|
| 109 |
+
|
| 110 |
+
def __init__(
|
| 111 |
+
self,
|
| 112 |
+
deps: ToolDependencies,
|
| 113 |
+
gradio_mode: bool = False,
|
| 114 |
+
instance_path: Optional[str] = None,
|
| 115 |
+
startup_voice: Optional[str] = None,
|
| 116 |
+
):
|
| 117 |
+
"""Initialize the handler."""
|
| 118 |
+
sample_rate = self.SAMPLE_RATE
|
| 119 |
+
super().__init__(
|
| 120 |
+
expected_layout="mono",
|
| 121 |
+
output_sample_rate=sample_rate,
|
| 122 |
+
input_sample_rate=sample_rate,
|
| 123 |
+
)
|
| 124 |
+
|
| 125 |
+
self.deps = deps
|
| 126 |
+
|
| 127 |
+
self.output_sample_rate = sample_rate
|
| 128 |
+
self.input_sample_rate = sample_rate
|
| 129 |
+
|
| 130 |
+
self.client: AsyncOpenAI
|
| 131 |
+
self.connection: AsyncRealtimeConnection | None = None
|
| 132 |
+
self.output_queue: "asyncio.Queue[Tuple[int, NDArray[np.int16]] | AdditionalOutputs]" = asyncio.Queue()
|
| 133 |
+
|
| 134 |
+
self.last_activity_time = asyncio.get_event_loop().time()
|
| 135 |
+
self.start_time = asyncio.get_event_loop().time()
|
| 136 |
+
self.is_idle_tool_call = False
|
| 137 |
+
self.gradio_mode = gradio_mode
|
| 138 |
+
self.instance_path = instance_path
|
| 139 |
+
self._voice_override: str | None = self._normalize_startup_voice(startup_voice)
|
| 140 |
+
self._realtime_connect_query: dict[str, str] = {}
|
| 141 |
+
|
| 142 |
+
# Debouncing for partial transcripts
|
| 143 |
+
self.partial_transcript_task: asyncio.Task[None] | None = None
|
| 144 |
+
self.partial_debounce_delay = 0.5 # seconds
|
| 145 |
+
self.input_transcript_chunks_by_item = InputTranscriptChunksByItem()
|
| 146 |
+
|
| 147 |
+
# Internal lifecycle flags
|
| 148 |
+
self._connected_event: asyncio.Event = asyncio.Event()
|
| 149 |
+
|
| 150 |
+
# Background tool manager
|
| 151 |
+
self.tool_manager = BackgroundToolManager()
|
| 152 |
+
|
| 153 |
+
# Cost tracking
|
| 154 |
+
self.cumulative_cost: float = 0.0
|
| 155 |
+
|
| 156 |
+
# Response-in-progress guard: the Realtime API only allows one active
|
| 157 |
+
# response per conversation at a time. A dedicated worker task
|
| 158 |
+
# (_response_sender_loop) dequeues and sends one request at a time
|
| 159 |
+
self._pending_responses: asyncio.Queue[dict[str, Any]] = asyncio.Queue()
|
| 160 |
+
self._response_done_event: asyncio.Event = asyncio.Event()
|
| 161 |
+
self._response_done_event.set()
|
| 162 |
+
self._response_started_or_rejected_event: asyncio.Event = asyncio.Event()
|
| 163 |
+
self._last_response_rejected: bool = False
|
| 164 |
+
self._turn_user_done_at: float | None = None
|
| 165 |
+
self._turn_response_created_at: float | None = None
|
| 166 |
+
self._turn_first_audio_at: float | None = None
|
| 167 |
+
|
| 168 |
+
@staticmethod
|
| 169 |
+
def _sanitize_tool_result_for_model(tool_name: str, tool_result: dict[str, Any]) -> dict[str, Any]:
|
| 170 |
+
"""Remove bulky transport-only fields before echoing tool output back to the model."""
|
| 171 |
+
if tool_name == "camera" and "b64_im" in tool_result:
|
| 172 |
+
sanitized = dict(tool_result)
|
| 173 |
+
sanitized.pop("b64_im", None)
|
| 174 |
+
sanitized["image_attached"] = True
|
| 175 |
+
return sanitized
|
| 176 |
+
return tool_result
|
| 177 |
+
|
| 178 |
+
def _normalize_startup_voice(self, voice: str | None) -> str | None:
|
| 179 |
+
"""Return a valid persisted startup voice for this backend, or None."""
|
| 180 |
+
available_voices = get_available_voices_for_backend(self.BACKEND_PROVIDER)
|
| 181 |
+
if voice in available_voices:
|
| 182 |
+
return voice
|
| 183 |
+
if voice:
|
| 184 |
+
logger.warning(
|
| 185 |
+
"Ignoring persisted startup voice %r for backend=%r; expected one of %s",
|
| 186 |
+
voice,
|
| 187 |
+
self.BACKEND_PROVIDER,
|
| 188 |
+
available_voices,
|
| 189 |
+
)
|
| 190 |
+
return None
|
| 191 |
+
|
| 192 |
+
def _response_done_timeout(self) -> float:
|
| 193 |
+
"""Return the response completion timeout."""
|
| 194 |
+
return _RESPONSE_DONE_TIMEOUT
|
| 195 |
+
|
| 196 |
+
def _connection_closed_errors(self) -> tuple[type[BaseException], ...]:
|
| 197 |
+
"""Return websocket closure exceptions handled as reconnectable/ignorable."""
|
| 198 |
+
return (ConnectionClosedError,)
|
| 199 |
+
|
| 200 |
+
@abstractmethod
|
| 201 |
+
def _get_session_instructions(self) -> str:
|
| 202 |
+
"""Return session instructions for this backend."""
|
| 203 |
+
|
| 204 |
+
@abstractmethod
|
| 205 |
+
def _get_session_voice(self, default: str | None = None) -> str:
|
| 206 |
+
"""Return the configured session voice for this backend."""
|
| 207 |
+
|
| 208 |
+
@abstractmethod
|
| 209 |
+
def _get_active_tool_specs(self) -> list[dict[str, Any]]:
|
| 210 |
+
"""Return active tool specs for the current session dependencies."""
|
| 211 |
+
|
| 212 |
+
@abstractmethod
|
| 213 |
+
def _get_session_config(self, tool_specs: list[dict[str, Any]]) -> RealtimeSessionCreateRequestParam:
|
| 214 |
+
"""Return the backend-specific realtime session config."""
|
| 215 |
+
|
| 216 |
+
async def _wait_for_output_item(self) -> Tuple[int, NDArray[np.int16]] | AdditionalOutputs | None:
|
| 217 |
+
"""Wait for the next output item."""
|
| 218 |
+
return await wait_for_item(self.output_queue) # type: ignore[no-any-return]
|
| 219 |
+
|
| 220 |
+
def _mark_activity(self, reason: str) -> None:
|
| 221 |
+
"""Record non-idle conversation activity for the idle timer."""
|
| 222 |
+
self.last_activity_time = asyncio.get_event_loop().time()
|
| 223 |
+
logger.debug("last activity time updated to %s (%s)", self.last_activity_time, reason)
|
| 224 |
+
|
| 225 |
+
def copy(self) -> "BaseRealtimeHandler":
|
| 226 |
+
"""Create a copy of the handler."""
|
| 227 |
+
return type(self)(
|
| 228 |
+
self.deps,
|
| 229 |
+
self.gradio_mode,
|
| 230 |
+
self.instance_path,
|
| 231 |
+
startup_voice=self._voice_override,
|
| 232 |
+
)
|
| 233 |
+
|
| 234 |
+
async def change_voice(self, voice: str) -> str:
|
| 235 |
+
"""Change only the voice and restart the session."""
|
| 236 |
+
self._voice_override = voice
|
| 237 |
+
if getattr(self, "client", None) is not None:
|
| 238 |
+
try:
|
| 239 |
+
await self._restart_session()
|
| 240 |
+
return f"Voice changed to {voice}."
|
| 241 |
+
except Exception as e:
|
| 242 |
+
logger.warning("Failed to restart session for voice change: %s", e)
|
| 243 |
+
return "Voice change failed. Will take effect on next connection."
|
| 244 |
+
return "Voice changed. Will take effect on next connection."
|
| 245 |
+
|
| 246 |
+
def get_current_voice(self) -> str:
|
| 247 |
+
"""Return the voice currently selected for this handler."""
|
| 248 |
+
default_voice = get_default_voice_for_backend(self.BACKEND_PROVIDER)
|
| 249 |
+
return self._voice_override or self._get_session_voice(default=default_voice)
|
| 250 |
+
|
| 251 |
+
async def apply_personality(self, profile: str | None) -> str:
|
| 252 |
+
"""Apply a new personality (profile) at runtime if possible.
|
| 253 |
+
|
| 254 |
+
- Updates the global config's selected profile for subsequent calls.
|
| 255 |
+
- If a realtime connection is active, sends a session.update with the
|
| 256 |
+
freshly resolved instructions so the change takes effect immediately.
|
| 257 |
+
|
| 258 |
+
Returns a short status message for UI feedback.
|
| 259 |
+
"""
|
| 260 |
+
try:
|
| 261 |
+
# Update the in-process config value and env
|
| 262 |
+
from reachy_mini_conversation_app.config import config as _config
|
| 263 |
+
from reachy_mini_conversation_app.config import set_custom_profile
|
| 264 |
+
|
| 265 |
+
set_custom_profile(profile)
|
| 266 |
+
logger.info(
|
| 267 |
+
"Set custom profile to %r (config=%r)", profile, getattr(_config, "REACHY_MINI_CUSTOM_PROFILE", None)
|
| 268 |
+
)
|
| 269 |
+
|
| 270 |
+
try:
|
| 271 |
+
instructions = self._get_session_instructions()
|
| 272 |
+
voice = self.get_current_voice()
|
| 273 |
+
except BaseException as e: # catch SystemExit from prompt loader without crashing
|
| 274 |
+
logger.error("Failed to resolve personality content: %s", e)
|
| 275 |
+
return f"Failed to apply personality: {e}"
|
| 276 |
+
|
| 277 |
+
# Attempt a live update first, then force a full restart to ensure it sticks
|
| 278 |
+
if self.connection is not None:
|
| 279 |
+
try:
|
| 280 |
+
await self.connection.session.update(
|
| 281 |
+
session=RealtimeSessionCreateRequestParam(
|
| 282 |
+
type="realtime",
|
| 283 |
+
instructions=instructions,
|
| 284 |
+
audio=RealtimeAudioConfigParam(
|
| 285 |
+
output=RealtimeAudioConfigOutputParam(
|
| 286 |
+
voice=voice,
|
| 287 |
+
),
|
| 288 |
+
),
|
| 289 |
+
),
|
| 290 |
+
)
|
| 291 |
+
logger.info("Applied personality via live update: %s", profile or "built-in default")
|
| 292 |
+
except Exception as e:
|
| 293 |
+
logger.warning("Live update failed; will restart session: %s", e)
|
| 294 |
+
|
| 295 |
+
# Force a real restart to guarantee the new instructions/voice
|
| 296 |
+
try:
|
| 297 |
+
await self._restart_session()
|
| 298 |
+
return "Applied personality and restarted realtime session."
|
| 299 |
+
except Exception as e:
|
| 300 |
+
logger.warning("Failed to restart session after apply: %s", e)
|
| 301 |
+
return "Applied personality. Will take effect on next connection."
|
| 302 |
+
else:
|
| 303 |
+
logger.info(
|
| 304 |
+
"Applied personality recorded: %s (no live connection; will apply on next session)",
|
| 305 |
+
profile or "built-in default",
|
| 306 |
+
)
|
| 307 |
+
return "Applied personality. Will take effect on next connection."
|
| 308 |
+
except Exception as e:
|
| 309 |
+
logger.error("Error applying personality '%s': %s", profile, e)
|
| 310 |
+
return f"Failed to apply personality: {e}"
|
| 311 |
+
|
| 312 |
+
async def _emit_debounced_partial(self, transcript: str, item_id: str, sequence_counter: int) -> None:
|
| 313 |
+
"""Emit partial transcript after debounce delay."""
|
| 314 |
+
try:
|
| 315 |
+
await asyncio.sleep(self.partial_debounce_delay)
|
| 316 |
+
|
| 317 |
+
input_transcript = self.input_transcript_chunks_by_item
|
| 318 |
+
if input_transcript.item_id == item_id and len(input_transcript.deltas) - 1 == sequence_counter:
|
| 319 |
+
await self.output_queue.put(AdditionalOutputs({"role": "user_partial", "content": transcript}))
|
| 320 |
+
logger.debug(f"Debounced partial emitted: {transcript}")
|
| 321 |
+
except asyncio.CancelledError:
|
| 322 |
+
logger.debug("Debounced partial cancelled")
|
| 323 |
+
raise
|
| 324 |
+
|
| 325 |
+
def _record_partial_transcript_delta(
|
| 326 |
+
self,
|
| 327 |
+
input_transcript: InputTranscriptChunksByItem,
|
| 328 |
+
item_id: str,
|
| 329 |
+
delta: str,
|
| 330 |
+
) -> None:
|
| 331 |
+
"""Record a suffix delta for a partial transcript."""
|
| 332 |
+
if input_transcript.item_id != item_id:
|
| 333 |
+
input_transcript.item_id = item_id
|
| 334 |
+
input_transcript.deltas = [delta]
|
| 335 |
+
else:
|
| 336 |
+
input_transcript.deltas.append(delta)
|
| 337 |
+
|
| 338 |
+
def _compute_response_cost(self, usage: Any) -> float:
|
| 339 |
+
"""Compute response cost using this backend's pricing."""
|
| 340 |
+
inp = getattr(usage, "input_token_details", None)
|
| 341 |
+
out = getattr(usage, "output_token_details", None)
|
| 342 |
+
cost = 0.0
|
| 343 |
+
if inp:
|
| 344 |
+
cost += (getattr(inp, "audio_tokens", 0) or 0) * self.AUDIO_INPUT_COST_PER_1M / 1e6
|
| 345 |
+
cost += (getattr(inp, "text_tokens", 0) or 0) * self.TEXT_INPUT_COST_PER_1M / 1e6
|
| 346 |
+
cost += (getattr(inp, "image_tokens", 0) or 0) * self.IMAGE_INPUT_COST_PER_1M / 1e6
|
| 347 |
+
if out:
|
| 348 |
+
cost += (getattr(out, "audio_tokens", 0) or 0) * self.AUDIO_OUTPUT_COST_PER_1M / 1e6
|
| 349 |
+
cost += (getattr(out, "text_tokens", 0) or 0) * self.TEXT_OUTPUT_COST_PER_1M / 1e6
|
| 350 |
+
return cost
|
| 351 |
+
|
| 352 |
+
async def _prepare_startup_credentials(self) -> None:
|
| 353 |
+
"""Let providers collect any startup credentials they need."""
|
| 354 |
+
|
| 355 |
+
def _persist_credentials_if_needed(self) -> None:
|
| 356 |
+
"""Let providers persist credentials after a successful session update."""
|
| 357 |
+
|
| 358 |
+
async def start_up(self) -> None:
|
| 359 |
+
"""Start the handler with minimal retries on unexpected websocket closure."""
|
| 360 |
+
await self._prepare_startup_credentials()
|
| 361 |
+
self.client = await self._build_realtime_client()
|
| 362 |
+
|
| 363 |
+
max_attempts = 3
|
| 364 |
+
for attempt in range(1, max_attempts + 1):
|
| 365 |
+
try:
|
| 366 |
+
await self._run_realtime_session()
|
| 367 |
+
# Normal exit from the session, stop retrying
|
| 368 |
+
return
|
| 369 |
+
except self._connection_closed_errors() as e:
|
| 370 |
+
# Abrupt close (e.g., "no close frame received or sent") → retry
|
| 371 |
+
logger.warning("Realtime websocket closed unexpectedly (attempt %d/%d): %s", attempt, max_attempts, e)
|
| 372 |
+
if attempt < max_attempts:
|
| 373 |
+
if self.REFRESH_CLIENT_ON_RECONNECT:
|
| 374 |
+
self.client = await self._build_realtime_client()
|
| 375 |
+
# exponential backoff with jitter
|
| 376 |
+
base_delay = 2 ** (attempt - 1) # 1s, 2s, 4s, 8s, etc.
|
| 377 |
+
jitter = random.uniform(0, 0.5)
|
| 378 |
+
delay = base_delay + jitter
|
| 379 |
+
logger.info("Retrying in %.1f seconds...", delay)
|
| 380 |
+
await asyncio.sleep(delay)
|
| 381 |
+
continue
|
| 382 |
+
raise
|
| 383 |
+
finally:
|
| 384 |
+
# never keep a stale reference
|
| 385 |
+
self.connection = None
|
| 386 |
+
try:
|
| 387 |
+
self._connected_event.clear()
|
| 388 |
+
except Exception:
|
| 389 |
+
pass
|
| 390 |
+
|
| 391 |
+
async def _restart_session(self) -> None:
|
| 392 |
+
"""Force-close the current session and start a fresh one in background.
|
| 393 |
+
|
| 394 |
+
Does not block the caller while the new session is establishing.
|
| 395 |
+
"""
|
| 396 |
+
try:
|
| 397 |
+
if self.connection is not None:
|
| 398 |
+
try:
|
| 399 |
+
await self.connection.close()
|
| 400 |
+
except Exception:
|
| 401 |
+
pass
|
| 402 |
+
finally:
|
| 403 |
+
self.connection = None
|
| 404 |
+
|
| 405 |
+
# Ensure we have a client (start_up must have run once)
|
| 406 |
+
if getattr(self, "client", None) is None:
|
| 407 |
+
logger.warning("Cannot restart: realtime client not initialized yet.")
|
| 408 |
+
return
|
| 409 |
+
|
| 410 |
+
# Fire-and-forget new session and wait briefly for connection
|
| 411 |
+
try:
|
| 412 |
+
self._connected_event.clear()
|
| 413 |
+
except Exception:
|
| 414 |
+
pass
|
| 415 |
+
if self.REFRESH_CLIENT_ON_RECONNECT:
|
| 416 |
+
self.client = await self._build_realtime_client()
|
| 417 |
+
asyncio.create_task(self._run_realtime_session(), name="realtime-session-restart")
|
| 418 |
+
try:
|
| 419 |
+
await asyncio.wait_for(self._connected_event.wait(), timeout=5.0)
|
| 420 |
+
logger.info("Realtime session restarted and connected.")
|
| 421 |
+
except asyncio.TimeoutError:
|
| 422 |
+
logger.warning("Realtime session restart timed out; continuing in background.")
|
| 423 |
+
except Exception as e:
|
| 424 |
+
logger.warning("_restart_session failed: %s", e)
|
| 425 |
+
|
| 426 |
+
async def _safe_response_create(self, **kwargs: Any) -> None:
|
| 427 |
+
"""Enqueue a response.create() kwargs for the sender worker _response_sender_loop().
|
| 428 |
+
|
| 429 |
+
This method never blocks the caller.
|
| 430 |
+
"""
|
| 431 |
+
await self._pending_responses.put(kwargs)
|
| 432 |
+
|
| 433 |
+
async def _response_sender_loop(self) -> None:
|
| 434 |
+
"""Dedicated worker that sends ``response.create()`` calls serially.
|
| 435 |
+
|
| 436 |
+
This logic was designed to comply with the response.create() docstring specification for event ordering:
|
| 437 |
+
https://github.com/openai/openai-python/blob/3e0c05b84a2056870abf3bd6a5e7849020209cc3/src/openai/resources/realtime/realtime.py#L649C1-L651C30
|
| 438 |
+
|
| 439 |
+
For each queued request the worker:
|
| 440 |
+
1. Waits until no response is active (_response_done_event).
|
| 441 |
+
2. Sends response.create().
|
| 442 |
+
3. Waits until the receiver observes response.created or a rejection.
|
| 443 |
+
4. Waits for the response cycle to complete (response.done).
|
| 444 |
+
5. If the server rejected with active_response, retries from step 1.
|
| 445 |
+
"""
|
| 446 |
+
while self.connection:
|
| 447 |
+
try:
|
| 448 |
+
kwargs = await self._pending_responses.get()
|
| 449 |
+
except asyncio.CancelledError:
|
| 450 |
+
return
|
| 451 |
+
|
| 452 |
+
sent = False
|
| 453 |
+
max_retries = 5
|
| 454 |
+
attempts = 0
|
| 455 |
+
while not sent and self.connection and attempts < max_retries:
|
| 456 |
+
try:
|
| 457 |
+
await asyncio.wait_for(
|
| 458 |
+
self._response_done_event.wait(),
|
| 459 |
+
timeout=self._response_done_timeout(),
|
| 460 |
+
)
|
| 461 |
+
except asyncio.TimeoutError:
|
| 462 |
+
logger.debug("Timed out waiting for previous response to finish; forcing ahead")
|
| 463 |
+
self._response_done_event.set()
|
| 464 |
+
|
| 465 |
+
if not self.connection:
|
| 466 |
+
break
|
| 467 |
+
|
| 468 |
+
self._last_response_rejected = False
|
| 469 |
+
self._response_started_or_rejected_event.clear()
|
| 470 |
+
try:
|
| 471 |
+
await self.connection.response.create(**kwargs)
|
| 472 |
+
except Exception as e:
|
| 473 |
+
logger.debug("_response_sender_loop: send failed: %s", e)
|
| 474 |
+
self._response_done_event.set()
|
| 475 |
+
break
|
| 476 |
+
|
| 477 |
+
try:
|
| 478 |
+
await asyncio.wait_for(
|
| 479 |
+
self._response_started_or_rejected_event.wait(),
|
| 480 |
+
timeout=self._response_done_timeout(),
|
| 481 |
+
)
|
| 482 |
+
except asyncio.TimeoutError:
|
| 483 |
+
logger.debug("Timed out waiting for response.created or response rejection")
|
| 484 |
+
|
| 485 |
+
# Check if the receiver loop observed an asynchronous rejection.
|
| 486 |
+
if self._last_response_rejected:
|
| 487 |
+
attempts += 1
|
| 488 |
+
if attempts >= max_retries:
|
| 489 |
+
logger.debug("response.create rejected %d times; giving up", attempts)
|
| 490 |
+
break
|
| 491 |
+
logger.debug("response.create was rejected; retrying (%d/%d)", attempts, max_retries)
|
| 492 |
+
await asyncio.sleep(_RESPONSE_REJECTION_RETRY_DELAY)
|
| 493 |
+
continue
|
| 494 |
+
|
| 495 |
+
try:
|
| 496 |
+
await asyncio.wait_for(
|
| 497 |
+
self._response_done_event.wait(),
|
| 498 |
+
timeout=self._response_done_timeout(),
|
| 499 |
+
)
|
| 500 |
+
except asyncio.TimeoutError:
|
| 501 |
+
logger.debug("Timed out waiting for response.done; assuming response completed")
|
| 502 |
+
self._response_done_event.set()
|
| 503 |
+
break
|
| 504 |
+
|
| 505 |
+
sent = True
|
| 506 |
+
|
| 507 |
+
async def _handle_tool_result(self, bg_tool: ToolNotification) -> None:
|
| 508 |
+
"""Process the result of a tool call."""
|
| 509 |
+
if bg_tool.error is not None:
|
| 510 |
+
logger.error("Tool '%s' (id=%s) failed with error: %s", bg_tool.tool_name, bg_tool.id, bg_tool.error)
|
| 511 |
+
tool_result = {"error": bg_tool.error}
|
| 512 |
+
tool_result_for_model = tool_result
|
| 513 |
+
elif bg_tool.result is not None:
|
| 514 |
+
tool_result = bg_tool.result
|
| 515 |
+
tool_result_for_model = (
|
| 516 |
+
self._sanitize_tool_result_for_model(bg_tool.tool_name, tool_result)
|
| 517 |
+
if isinstance(tool_result, dict)
|
| 518 |
+
else tool_result
|
| 519 |
+
)
|
| 520 |
+
logger.info(
|
| 521 |
+
"Tool '%s' (id=%s) executed successfully.",
|
| 522 |
+
bg_tool.tool_name,
|
| 523 |
+
bg_tool.id,
|
| 524 |
+
)
|
| 525 |
+
logger.debug("Tool '%s' model-visible result: %s", bg_tool.tool_name, tool_result_for_model)
|
| 526 |
+
else:
|
| 527 |
+
logger.warning("Tool '%s' (id=%s) returned no result and no error", bg_tool.tool_name, bg_tool.id)
|
| 528 |
+
tool_result = {"error": "No result returned from tool execution"}
|
| 529 |
+
tool_result_for_model = tool_result
|
| 530 |
+
|
| 531 |
+
# Connection may have closed while tool was running
|
| 532 |
+
if not self.connection:
|
| 533 |
+
logger.warning(
|
| 534 |
+
"Connection closed during tool '%s' (id=%s) execution; cannot send result back",
|
| 535 |
+
bg_tool.tool_name,
|
| 536 |
+
bg_tool.id,
|
| 537 |
+
)
|
| 538 |
+
return
|
| 539 |
+
|
| 540 |
+
try:
|
| 541 |
+
self._mark_activity("tool_result_ready")
|
| 542 |
+
if isinstance(bg_tool.id, str):
|
| 543 |
+
await self.connection.conversation.item.create(
|
| 544 |
+
item={
|
| 545 |
+
"type": "function_call_output",
|
| 546 |
+
"call_id": bg_tool.id,
|
| 547 |
+
"output": json.dumps(tool_result_for_model),
|
| 548 |
+
},
|
| 549 |
+
)
|
| 550 |
+
|
| 551 |
+
await self.output_queue.put(
|
| 552 |
+
AdditionalOutputs(
|
| 553 |
+
{
|
| 554 |
+
"role": "assistant",
|
| 555 |
+
"content": json.dumps(tool_result_for_model),
|
| 556 |
+
# Gradio UI metadata.status accept only "pending" and "done". Do not accept bg.tool.status values.
|
| 557 |
+
"metadata": {
|
| 558 |
+
"title": f"🛠️ Used tool {bg_tool.tool_name}",
|
| 559 |
+
"status": "done",
|
| 560 |
+
},
|
| 561 |
+
},
|
| 562 |
+
),
|
| 563 |
+
)
|
| 564 |
+
|
| 565 |
+
if bg_tool.tool_name == "camera" and "b64_im" in tool_result:
|
| 566 |
+
# use raw base64, don't json.dumps (which adds quotes)
|
| 567 |
+
b64_im = tool_result["b64_im"]
|
| 568 |
+
if not isinstance(b64_im, str):
|
| 569 |
+
logger.warning("Unexpected type for b64_im: %s", type(b64_im))
|
| 570 |
+
b64_im = str(b64_im)
|
| 571 |
+
image_width = tool_result.get("image_width")
|
| 572 |
+
image_height = tool_result.get("image_height")
|
| 573 |
+
jpeg_bytes_value = tool_result.get("jpeg_bytes")
|
| 574 |
+
jpeg_bytes = jpeg_bytes_value if isinstance(jpeg_bytes_value, int) else (len(b64_im) * 3) // 4
|
| 575 |
+
await self.connection.conversation.item.create(
|
| 576 |
+
item={
|
| 577 |
+
"type": "message",
|
| 578 |
+
"role": "user",
|
| 579 |
+
"content": [
|
| 580 |
+
{
|
| 581 |
+
"type": "input_image",
|
| 582 |
+
"image_url": f"data:image/jpeg;base64,{b64_im}",
|
| 583 |
+
},
|
| 584 |
+
],
|
| 585 |
+
},
|
| 586 |
+
)
|
| 587 |
+
if isinstance(image_width, int) and isinstance(image_height, int):
|
| 588 |
+
logger.info(
|
| 589 |
+
"Added camera image to conversation frame=%sx%s jpeg_bytes=%s",
|
| 590 |
+
image_width,
|
| 591 |
+
image_height,
|
| 592 |
+
jpeg_bytes,
|
| 593 |
+
)
|
| 594 |
+
else:
|
| 595 |
+
logger.info(
|
| 596 |
+
"Added camera image to conversation jpeg_bytes=%s",
|
| 597 |
+
jpeg_bytes,
|
| 598 |
+
)
|
| 599 |
+
|
| 600 |
+
if self.deps.camera_worker is not None:
|
| 601 |
+
np_img = self.deps.camera_worker.get_latest_frame()
|
| 602 |
+
if np_img is not None:
|
| 603 |
+
# Camera frames are BGR; reverse channels without requiring OpenCV in core installs.
|
| 604 |
+
rgb_frame = np_img[:, :, ::-1].copy() if np_img.ndim == 3 and np_img.shape[-1] == 3 else np_img
|
| 605 |
+
else:
|
| 606 |
+
rgb_frame = None
|
| 607 |
+
img = gr.Image(value=rgb_frame)
|
| 608 |
+
|
| 609 |
+
await self.output_queue.put(
|
| 610 |
+
AdditionalOutputs(
|
| 611 |
+
{
|
| 612 |
+
"role": "assistant",
|
| 613 |
+
"content": img,
|
| 614 |
+
},
|
| 615 |
+
),
|
| 616 |
+
)
|
| 617 |
+
|
| 618 |
+
# If this tool call was triggered by an idle signal, don't make the robot speak.
|
| 619 |
+
# For other tool calls, let the robot reply out loud.
|
| 620 |
+
if not bg_tool.is_idle_tool_call:
|
| 621 |
+
await self._safe_response_create(
|
| 622 |
+
response=RealtimeResponseCreateParamsParam(
|
| 623 |
+
instructions="Use the tool result just returned and answer concisely in speech.",
|
| 624 |
+
),
|
| 625 |
+
)
|
| 626 |
+
|
| 627 |
+
except self._connection_closed_errors():
|
| 628 |
+
logger.warning("Connection closed while sending tool result")
|
| 629 |
+
self.connection = None
|
| 630 |
+
self._response_done_event.set()
|
| 631 |
+
|
| 632 |
+
async def _run_realtime_session(self) -> None:
|
| 633 |
+
"""Establish and manage a single realtime session."""
|
| 634 |
+
tool_specs = self._get_active_tool_specs()
|
| 635 |
+
logger.info(
|
| 636 |
+
"Tools to be used in conversation: %s",
|
| 637 |
+
[tool["name"] for tool in tool_specs],
|
| 638 |
+
)
|
| 639 |
+
connect_kwargs: dict[str, Any] = {}
|
| 640 |
+
if config.MODEL_NAME:
|
| 641 |
+
connect_kwargs["model"] = config.MODEL_NAME
|
| 642 |
+
if self._realtime_connect_query:
|
| 643 |
+
connect_kwargs["extra_query"] = self._realtime_connect_query
|
| 644 |
+
async with self.client.realtime.connect(**connect_kwargs) as conn:
|
| 645 |
+
try:
|
| 646 |
+
session_config = self._get_session_config(tool_specs)
|
| 647 |
+
await conn.session.update(session=session_config)
|
| 648 |
+
logger.info(
|
| 649 |
+
"Realtime session initialized with profile=%r voice=%r",
|
| 650 |
+
getattr(config, "REACHY_MINI_CUSTOM_PROFILE", None),
|
| 651 |
+
self.get_current_voice(),
|
| 652 |
+
)
|
| 653 |
+
self._persist_credentials_if_needed()
|
| 654 |
+
except Exception:
|
| 655 |
+
logger.exception("Realtime session.update failed; aborting startup")
|
| 656 |
+
raise
|
| 657 |
+
|
| 658 |
+
logger.info("Realtime session updated successfully")
|
| 659 |
+
|
| 660 |
+
# Reset the partial-transcript accumulator for each new session
|
| 661 |
+
self.input_transcript_chunks_by_item = InputTranscriptChunksByItem()
|
| 662 |
+
|
| 663 |
+
# Manage events received from the realtime server.
|
| 664 |
+
self.connection = conn
|
| 665 |
+
try:
|
| 666 |
+
self._connected_event.set()
|
| 667 |
+
except Exception:
|
| 668 |
+
pass
|
| 669 |
+
|
| 670 |
+
response_sender_task: asyncio.Task[None] | None = None
|
| 671 |
+
try:
|
| 672 |
+
# Start the background tool manager
|
| 673 |
+
self.tool_manager.start_up(tool_callbacks=[self._handle_tool_result])
|
| 674 |
+
|
| 675 |
+
# Start the response sender worker
|
| 676 |
+
response_sender_task = asyncio.create_task(self._response_sender_loop(), name="response-sender")
|
| 677 |
+
|
| 678 |
+
async for event in self.connection:
|
| 679 |
+
logger.debug("Realtime event: %s", event.type)
|
| 680 |
+
if event.type == "input_audio_buffer.speech_started":
|
| 681 |
+
self._mark_activity("user_speech_started")
|
| 682 |
+
self._turn_user_done_at = None
|
| 683 |
+
self._turn_response_created_at = None
|
| 684 |
+
self._turn_first_audio_at = None
|
| 685 |
+
if hasattr(self, "_clear_queue") and callable(self._clear_queue):
|
| 686 |
+
self._clear_queue()
|
| 687 |
+
if self.deps.head_wobbler is not None:
|
| 688 |
+
self.deps.head_wobbler.reset()
|
| 689 |
+
self.deps.movement_manager.set_listening(True)
|
| 690 |
+
logger.debug("User speech started")
|
| 691 |
+
|
| 692 |
+
if event.type == "input_audio_buffer.speech_stopped":
|
| 693 |
+
self._mark_activity("user_speech_stopped")
|
| 694 |
+
self.deps.movement_manager.set_listening(False)
|
| 695 |
+
logger.debug("User speech stopped - server will auto-commit with VAD")
|
| 696 |
+
|
| 697 |
+
if event.type == "response.output_audio.done":
|
| 698 |
+
if self.deps.head_wobbler is not None:
|
| 699 |
+
self.deps.head_wobbler.request_reset_after_current_audio()
|
| 700 |
+
logger.debug("response completed")
|
| 701 |
+
|
| 702 |
+
if event.type == "response.created":
|
| 703 |
+
self._mark_activity("response_created")
|
| 704 |
+
self._response_done_event.clear()
|
| 705 |
+
self._response_started_or_rejected_event.set()
|
| 706 |
+
if self._turn_user_done_at is not None and self._turn_response_created_at is None:
|
| 707 |
+
self._turn_response_created_at = time.perf_counter()
|
| 708 |
+
delta_ms = (self._turn_response_created_at - self._turn_user_done_at) * 1000
|
| 709 |
+
logger.info("Turn latency: response.created %.0f ms after user transcript", delta_ms)
|
| 710 |
+
logger.debug("Response created (active)")
|
| 711 |
+
|
| 712 |
+
if event.type == "response.done":
|
| 713 |
+
# Doesn't mean the audio is done playing
|
| 714 |
+
self._response_done_event.set()
|
| 715 |
+
self._response_started_or_rejected_event.set()
|
| 716 |
+
self.is_idle_tool_call = False
|
| 717 |
+
logger.debug("Response done")
|
| 718 |
+
|
| 719 |
+
response = getattr(event, "response", None)
|
| 720 |
+
usage = getattr(response, "usage", None) if response else None
|
| 721 |
+
if usage:
|
| 722 |
+
cost = self._compute_response_cost(usage)
|
| 723 |
+
self.cumulative_cost += cost
|
| 724 |
+
logger.debug("Cost: $%.4f | Cumulative: $%.4f", cost, self.cumulative_cost)
|
| 725 |
+
else:
|
| 726 |
+
logger.warning("No usage data available for cost tracking")
|
| 727 |
+
|
| 728 |
+
if event.type == "conversation.item.input_audio_transcription.delta":
|
| 729 |
+
self._mark_activity("user_transcription_delta")
|
| 730 |
+
logger.debug(f"User partial transcript: {event.delta}")
|
| 731 |
+
|
| 732 |
+
item_id = event.item_id
|
| 733 |
+
delta = event.delta or ""
|
| 734 |
+
|
| 735 |
+
input_transcript = self.input_transcript_chunks_by_item
|
| 736 |
+
self._record_partial_transcript_delta(input_transcript, item_id, delta)
|
| 737 |
+
|
| 738 |
+
current_partial = "".join(input_transcript.deltas)
|
| 739 |
+
sequence_counter = len(input_transcript.deltas) - 1
|
| 740 |
+
|
| 741 |
+
# Cancel previous debounce task if it exists
|
| 742 |
+
if self.partial_transcript_task and not self.partial_transcript_task.done():
|
| 743 |
+
self.partial_transcript_task.cancel()
|
| 744 |
+
try:
|
| 745 |
+
await self.partial_transcript_task
|
| 746 |
+
except asyncio.CancelledError:
|
| 747 |
+
pass
|
| 748 |
+
|
| 749 |
+
# Start new debounce timer with the last delta
|
| 750 |
+
self.partial_transcript_task = asyncio.create_task(
|
| 751 |
+
self._emit_debounced_partial(current_partial, item_id, sequence_counter)
|
| 752 |
+
)
|
| 753 |
+
|
| 754 |
+
# Handle completed transcription (user finished speaking)
|
| 755 |
+
if event.type == "conversation.item.input_audio_transcription.completed":
|
| 756 |
+
self._mark_activity("user_transcription_completed")
|
| 757 |
+
raw_transcript = event.transcript or ""
|
| 758 |
+
transcript = raw_transcript.strip()
|
| 759 |
+
logger.debug("User transcript: %s", raw_transcript)
|
| 760 |
+
self.deps.movement_manager.set_listening(False)
|
| 761 |
+
|
| 762 |
+
# Cancel any pending partial emission
|
| 763 |
+
if self.partial_transcript_task and not self.partial_transcript_task.done():
|
| 764 |
+
self.partial_transcript_task.cancel()
|
| 765 |
+
try:
|
| 766 |
+
await self.partial_transcript_task
|
| 767 |
+
except asyncio.CancelledError:
|
| 768 |
+
pass
|
| 769 |
+
|
| 770 |
+
if not transcript:
|
| 771 |
+
logger.debug("Ignoring empty user transcript")
|
| 772 |
+
continue
|
| 773 |
+
|
| 774 |
+
self._turn_user_done_at = time.perf_counter()
|
| 775 |
+
self._turn_response_created_at = None
|
| 776 |
+
self._turn_first_audio_at = None
|
| 777 |
+
|
| 778 |
+
await self.output_queue.put(AdditionalOutputs({"role": "user", "content": transcript}))
|
| 779 |
+
|
| 780 |
+
# Handle assistant transcription
|
| 781 |
+
if event.type == "response.output_audio_transcript.done":
|
| 782 |
+
self._mark_activity("assistant_transcript_done")
|
| 783 |
+
logger.debug(f"Assistant transcript: {event.transcript}")
|
| 784 |
+
await self.output_queue.put(
|
| 785 |
+
AdditionalOutputs({"role": "assistant", "content": event.transcript})
|
| 786 |
+
)
|
| 787 |
+
|
| 788 |
+
# Handle audio delta
|
| 789 |
+
if event.type == "response.output_audio.delta":
|
| 790 |
+
decoded_pcm_bytes = base64.b64decode(event.delta)
|
| 791 |
+
decoded_pcm = np.frombuffer(decoded_pcm_bytes, dtype=np.int16).reshape(1, -1)
|
| 792 |
+
if self.gradio_mode and self.deps.head_wobbler is not None:
|
| 793 |
+
self.deps.head_wobbler.feed_pcm(decoded_pcm, self.output_sample_rate)
|
| 794 |
+
self._mark_activity("assistant_audio_delta")
|
| 795 |
+
if self._turn_user_done_at is not None and self._turn_first_audio_at is None:
|
| 796 |
+
self._turn_first_audio_at = time.perf_counter()
|
| 797 |
+
delta_ms = (self._turn_first_audio_at - self._turn_user_done_at) * 1000
|
| 798 |
+
logger.info("Turn latency: first audio delta %.0f ms after user transcript", delta_ms)
|
| 799 |
+
await self.output_queue.put(
|
| 800 |
+
(
|
| 801 |
+
self.output_sample_rate,
|
| 802 |
+
decoded_pcm,
|
| 803 |
+
),
|
| 804 |
+
)
|
| 805 |
+
# ---- tool-calling plumbing ----
|
| 806 |
+
if event.type == "response.function_call_arguments.done":
|
| 807 |
+
self._mark_activity("tool_call_received")
|
| 808 |
+
tool_name = getattr(event, "name", None)
|
| 809 |
+
args_json_str = getattr(event, "arguments", None)
|
| 810 |
+
call_id: str = str(getattr(event, "call_id", uuid.uuid4()))
|
| 811 |
+
|
| 812 |
+
logger.info(
|
| 813 |
+
"Tool call received — tool_name=%r, call_id=%s, is_idle=%s, args=%s",
|
| 814 |
+
tool_name,
|
| 815 |
+
call_id,
|
| 816 |
+
self.is_idle_tool_call,
|
| 817 |
+
args_json_str,
|
| 818 |
+
)
|
| 819 |
+
|
| 820 |
+
if not isinstance(tool_name, str) or not isinstance(args_json_str, str):
|
| 821 |
+
logger.error(
|
| 822 |
+
"Invalid tool call: tool_name=%s (type=%s), args=%s (type=%s), call_id=%s",
|
| 823 |
+
tool_name,
|
| 824 |
+
type(tool_name).__name__,
|
| 825 |
+
args_json_str,
|
| 826 |
+
type(args_json_str).__name__,
|
| 827 |
+
call_id,
|
| 828 |
+
)
|
| 829 |
+
continue
|
| 830 |
+
|
| 831 |
+
bg_tool = await self.tool_manager.start_tool(
|
| 832 |
+
call_id=call_id,
|
| 833 |
+
tool_call_routine=ToolCallRoutine(
|
| 834 |
+
tool_name=tool_name,
|
| 835 |
+
args_json_str=args_json_str,
|
| 836 |
+
deps=self.deps,
|
| 837 |
+
),
|
| 838 |
+
is_idle_tool_call=self.is_idle_tool_call,
|
| 839 |
+
)
|
| 840 |
+
|
| 841 |
+
await self.output_queue.put(
|
| 842 |
+
AdditionalOutputs(
|
| 843 |
+
{
|
| 844 |
+
"role": "assistant",
|
| 845 |
+
"content": f"🛠️ Used tool {tool_name} with args {args_json_str}. The tool is now running. Tool ID: {bg_tool.tool_id}",
|
| 846 |
+
},
|
| 847 |
+
),
|
| 848 |
+
)
|
| 849 |
+
logger.info(
|
| 850 |
+
"Started background tool: %s (id=%s, call_id=%s)", tool_name, bg_tool.tool_id, call_id
|
| 851 |
+
)
|
| 852 |
+
|
| 853 |
+
# server error
|
| 854 |
+
if event.type == "error":
|
| 855 |
+
err = getattr(event, "error", None)
|
| 856 |
+
msg = getattr(err, "message", str(err) if err else "unknown error")
|
| 857 |
+
code = getattr(err, "code", "") or getattr(err, "type", "")
|
| 858 |
+
|
| 859 |
+
if code == "conversation_already_has_active_response":
|
| 860 |
+
# response.create was rejected. The sender worker
|
| 861 |
+
# is waiting on _response_done_event; when the active
|
| 862 |
+
# response finishes it will wake up and see this flag.
|
| 863 |
+
self._last_response_rejected = True
|
| 864 |
+
self._response_started_or_rejected_event.set()
|
| 865 |
+
logger.debug("response.create rejected; worker will retry after active response finishes")
|
| 866 |
+
else:
|
| 867 |
+
self._response_started_or_rejected_event.set()
|
| 868 |
+
logger.error("Realtime error [%s]: %s (raw=%s)", code, msg, err)
|
| 869 |
+
|
| 870 |
+
if code == "input_audio_buffer_commit_empty":
|
| 871 |
+
self.deps.movement_manager.set_listening(False)
|
| 872 |
+
|
| 873 |
+
# Only show user-facing errors, not internal state errors.
|
| 874 |
+
if code not in ("input_audio_buffer_commit_empty", "conversation_already_has_active_response"):
|
| 875 |
+
await self.output_queue.put(
|
| 876 |
+
AdditionalOutputs({"role": "assistant", "content": f"[error] {msg}"})
|
| 877 |
+
)
|
| 878 |
+
finally:
|
| 879 |
+
# Stop the response sender worker.
|
| 880 |
+
if response_sender_task is not None:
|
| 881 |
+
response_sender_task.cancel()
|
| 882 |
+
try:
|
| 883 |
+
await response_sender_task
|
| 884 |
+
except asyncio.CancelledError:
|
| 885 |
+
pass
|
| 886 |
+
|
| 887 |
+
# Stop background tool manager tasks (listener + cleanup) in all paths.
|
| 888 |
+
await self.tool_manager.shutdown()
|
| 889 |
+
|
| 890 |
+
# Microphone receive
|
| 891 |
+
async def receive(self, frame: Tuple[int, NDArray[np.int16]]) -> None:
|
| 892 |
+
"""Receive audio frame from the microphone and send it to the realtime server.
|
| 893 |
+
|
| 894 |
+
Handles both mono and stereo audio formats, converting to the expected
|
| 895 |
+
mono format for the realtime API. Resamples if the input sample rate differs
|
| 896 |
+
from the expected rate.
|
| 897 |
+
|
| 898 |
+
Args:
|
| 899 |
+
frame: A tuple containing (sample_rate, audio_data).
|
| 900 |
+
|
| 901 |
+
"""
|
| 902 |
+
if not self.connection:
|
| 903 |
+
return
|
| 904 |
+
|
| 905 |
+
input_sample_rate, audio_frame = frame
|
| 906 |
+
|
| 907 |
+
# Reshape if needed
|
| 908 |
+
if audio_frame.ndim == 2:
|
| 909 |
+
# Scipy channels last convention
|
| 910 |
+
if audio_frame.shape[1] > audio_frame.shape[0]:
|
| 911 |
+
audio_frame = audio_frame.T
|
| 912 |
+
# Multiple channels -> Mono channel
|
| 913 |
+
if audio_frame.shape[1] > 1:
|
| 914 |
+
audio_frame = audio_frame[:, 0]
|
| 915 |
+
|
| 916 |
+
# Resample if needed
|
| 917 |
+
if self.input_sample_rate != input_sample_rate:
|
| 918 |
+
audio_frame = resample(audio_frame, int(len(audio_frame) * self.input_sample_rate / input_sample_rate))
|
| 919 |
+
|
| 920 |
+
# Cast if needed
|
| 921 |
+
audio_frame = audio_to_int16(audio_frame)
|
| 922 |
+
|
| 923 |
+
# Send to the realtime input buffer (guard against races during reconnect).
|
| 924 |
+
try:
|
| 925 |
+
audio_message = base64.b64encode(audio_frame.tobytes()).decode("utf-8")
|
| 926 |
+
await self.connection.input_audio_buffer.append(audio=audio_message)
|
| 927 |
+
except Exception as e:
|
| 928 |
+
logger.debug("Dropping audio frame: connection not ready (%s)", e)
|
| 929 |
+
return
|
| 930 |
+
|
| 931 |
+
async def emit(self) -> Tuple[int, NDArray[np.int16]] | AdditionalOutputs | None:
|
| 932 |
+
"""Emit audio frame to be played by the speaker."""
|
| 933 |
+
# Sends output queued by the realtime event handler to the stream.
|
| 934 |
+
# This is called periodically by the fastrtc Stream
|
| 935 |
+
|
| 936 |
+
# Handle idle
|
| 937 |
+
idle_duration = asyncio.get_event_loop().time() - self.last_activity_time
|
| 938 |
+
if idle_duration > 180.0 and self._response_done_event.is_set() and self.deps.movement_manager.is_idle():
|
| 939 |
+
try:
|
| 940 |
+
await self.send_idle_signal(idle_duration)
|
| 941 |
+
except Exception as e:
|
| 942 |
+
logger.warning("Idle signal skipped (connection closed?): %s", e)
|
| 943 |
+
return None
|
| 944 |
+
|
| 945 |
+
self.last_activity_time = asyncio.get_event_loop().time() # avoid repeated resets
|
| 946 |
+
|
| 947 |
+
return await self._wait_for_output_item()
|
| 948 |
+
|
| 949 |
+
async def shutdown(self) -> None:
|
| 950 |
+
"""Shutdown the handler."""
|
| 951 |
+
# Unblock the response sender worker so it can exit
|
| 952 |
+
self._response_done_event.set()
|
| 953 |
+
|
| 954 |
+
# Stop background tool manager tasks (listener + cleanup)
|
| 955 |
+
await self.tool_manager.shutdown()
|
| 956 |
+
|
| 957 |
+
# Cancel any pending debounce task
|
| 958 |
+
if self.partial_transcript_task and not self.partial_transcript_task.done():
|
| 959 |
+
self.partial_transcript_task.cancel()
|
| 960 |
+
try:
|
| 961 |
+
await self.partial_transcript_task
|
| 962 |
+
except asyncio.CancelledError:
|
| 963 |
+
pass
|
| 964 |
+
|
| 965 |
+
if self.connection:
|
| 966 |
+
try:
|
| 967 |
+
await self.connection.close()
|
| 968 |
+
except self._connection_closed_errors() as e:
|
| 969 |
+
logger.debug(f"Connection already closed during shutdown: {e}")
|
| 970 |
+
except Exception as e:
|
| 971 |
+
logger.debug(f"connection.close() ignored: {e}")
|
| 972 |
+
finally:
|
| 973 |
+
self.connection = None
|
| 974 |
+
|
| 975 |
+
# Clear any remaining items in the output queue
|
| 976 |
+
while not self.output_queue.empty():
|
| 977 |
+
try:
|
| 978 |
+
self.output_queue.get_nowait()
|
| 979 |
+
except asyncio.QueueEmpty:
|
| 980 |
+
break
|
| 981 |
+
|
| 982 |
+
def format_timestamp(self) -> str:
|
| 983 |
+
"""Format current timestamp with date, time, and elapsed seconds."""
|
| 984 |
+
loop_time = asyncio.get_event_loop().time() # monotonic
|
| 985 |
+
elapsed_seconds = loop_time - self.start_time
|
| 986 |
+
dt = datetime.now() # wall-clock
|
| 987 |
+
return f"[{dt.strftime('%Y-%m-%d %H:%M:%S')} | +{elapsed_seconds:.1f}s]"
|
| 988 |
+
|
| 989 |
+
async def get_available_voices(self) -> list[str]:
|
| 990 |
+
"""Return available voices for this backend."""
|
| 991 |
+
return get_available_voices_for_backend(self.BACKEND_PROVIDER)
|
| 992 |
+
|
| 993 |
+
@abstractmethod
|
| 994 |
+
async def _build_realtime_client(self) -> AsyncOpenAI:
|
| 995 |
+
"""Build the realtime SDK client for this backend."""
|
| 996 |
+
|
| 997 |
+
async def send_idle_signal(self, idle_duration: float) -> None:
|
| 998 |
+
"""Send an idle signal to the realtime server."""
|
| 999 |
+
logger.debug("Sending idle signal")
|
| 1000 |
+
self.is_idle_tool_call = True
|
| 1001 |
+
timestamp_msg = f"[Idle time update: {self.format_timestamp()} - No activity for {idle_duration:.1f}s] You've been idle for a while. Feel free to get creative - dance, show an emotion, look around, call idle_do_nothing to stay still and silent, or just be yourself!"
|
| 1002 |
+
if not self.connection:
|
| 1003 |
+
logger.debug("No connection, cannot send idle signal")
|
| 1004 |
+
return
|
| 1005 |
+
await self.connection.conversation.item.create(
|
| 1006 |
+
item={
|
| 1007 |
+
"type": "message",
|
| 1008 |
+
"role": "user",
|
| 1009 |
+
"content": [{"type": "input_text", "text": timestamp_msg}],
|
| 1010 |
+
},
|
| 1011 |
+
)
|
| 1012 |
+
await self._safe_response_create(
|
| 1013 |
+
response=RealtimeResponseCreateParamsParam(
|
| 1014 |
+
instructions="You MUST respond with function calls only - no speech or text. Choose appropriate actions for idle behavior. Use idle_do_nothing only if you intentionally want no movement or sound during this idle turn.",
|
| 1015 |
+
tool_choice="required",
|
| 1016 |
+
),
|
| 1017 |
+
)
|
src/reachy_mini_conversation_app/config.py
CHANGED
|
@@ -2,6 +2,8 @@ import os
|
|
| 2 |
import sys
|
| 3 |
import logging
|
| 4 |
from pathlib import Path
|
|
|
|
|
|
|
| 5 |
from importlib.resources import files
|
| 6 |
|
| 7 |
from dotenv import find_dotenv, load_dotenv
|
|
@@ -56,6 +58,20 @@ AVAILABLE_VOICES: list[str] = [
|
|
| 56 |
"shimmer",
|
| 57 |
"verse",
|
| 58 |
]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 59 |
|
| 60 |
# Voices supported by the Gemini Live API
|
| 61 |
GEMINI_AVAILABLE_VOICES: list[str] = [
|
|
@@ -71,14 +87,45 @@ GEMINI_AVAILABLE_VOICES: list[str] = [
|
|
| 71 |
|
| 72 |
OPENAI_BACKEND = "openai"
|
| 73 |
GEMINI_BACKEND = "gemini"
|
| 74 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 75 |
DEFAULT_MODEL_NAME_BY_BACKEND = {
|
| 76 |
OPENAI_BACKEND: "gpt-realtime",
|
| 77 |
GEMINI_BACKEND: "gemini-3.1-flash-live-preview",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 78 |
}
|
| 79 |
DEFAULT_VOICE_BY_BACKEND = {
|
| 80 |
-
OPENAI_BACKEND:
|
| 81 |
GEMINI_BACKEND: "Kore",
|
|
|
|
| 82 |
}
|
| 83 |
|
| 84 |
logger = logging.getLogger(__name__)
|
|
@@ -94,10 +141,13 @@ def _normalize_backend_provider(
|
|
| 94 |
backend_provider: str | None = None,
|
| 95 |
model_name: str | None = None,
|
| 96 |
) -> str:
|
| 97 |
-
"""Normalize
|
| 98 |
candidate = (backend_provider or "").strip().lower()
|
| 99 |
if candidate in DEFAULT_MODEL_NAME_BY_BACKEND:
|
| 100 |
return candidate
|
|
|
|
|
|
|
|
|
|
| 101 |
return GEMINI_BACKEND if _is_gemini_model_name(model_name) else DEFAULT_BACKEND_PROVIDER
|
| 102 |
|
| 103 |
|
|
@@ -107,11 +157,14 @@ def _resolve_model_name(
|
|
| 107 |
) -> str:
|
| 108 |
"""Return a model name that matches the selected backend provider."""
|
| 109 |
normalized_backend = _normalize_backend_provider(backend_provider, model_name)
|
|
|
|
|
|
|
|
|
|
| 110 |
candidate = (model_name or "").strip()
|
| 111 |
if candidate:
|
| 112 |
if normalized_backend == GEMINI_BACKEND and _is_gemini_model_name(candidate):
|
| 113 |
return candidate
|
| 114 |
-
if normalized_backend =
|
| 115 |
return candidate
|
| 116 |
logger.warning(
|
| 117 |
"MODEL_NAME=%r does not match BACKEND_PROVIDER=%r, using default %r",
|
|
@@ -142,6 +195,92 @@ def _env_flag(name: str, default: bool = False) -> bool:
|
|
| 142 |
return default
|
| 143 |
|
| 144 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 145 |
def _collect_profile_names(profiles_root: Path) -> set[str]:
|
| 146 |
"""Return profile folder names from a profiles root directory."""
|
| 147 |
if not profiles_root.exists() or not profiles_root.is_dir():
|
|
@@ -218,14 +357,23 @@ class Config:
|
|
| 218 |
os.getenv("MODEL_NAME"),
|
| 219 |
)
|
| 220 |
MODEL_NAME = _resolve_model_name(BACKEND_PROVIDER, os.getenv("MODEL_NAME"))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 221 |
HF_HOME = os.getenv("HF_HOME", "./cache")
|
| 222 |
LOCAL_VISION_MODEL = os.getenv("LOCAL_VISION_MODEL", "HuggingFaceTB/SmolVLM2-2.2B-Instruct")
|
| 223 |
HF_TOKEN = os.getenv("HF_TOKEN") # Optional, falls back to hf auth login if not set
|
| 224 |
|
| 225 |
logger.debug(
|
| 226 |
-
"Backend provider: %s, Model: %s, HF_HOME: %s, Vision Model: %s",
|
| 227 |
BACKEND_PROVIDER,
|
| 228 |
MODEL_NAME,
|
|
|
|
|
|
|
|
|
|
| 229 |
HF_HOME,
|
| 230 |
LOCAL_VISION_MODEL,
|
| 231 |
)
|
|
@@ -312,6 +460,15 @@ def refresh_runtime_config_from_env() -> None:
|
|
| 312 |
os.getenv("MODEL_NAME"),
|
| 313 |
)
|
| 314 |
config.MODEL_NAME = _resolve_model_name(config.BACKEND_PROVIDER, os.getenv("MODEL_NAME"))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 315 |
config.REACHY_MINI_CUSTOM_PROFILE = LOCKED_PROFILE or os.getenv("REACHY_MINI_CUSTOM_PROFILE")
|
| 316 |
|
| 317 |
|
|
@@ -327,11 +484,19 @@ def get_model_name_for_backend(backend: str) -> str:
|
|
| 327 |
return DEFAULT_MODEL_NAME_BY_BACKEND[_normalize_backend_provider(backend)]
|
| 328 |
|
| 329 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 330 |
def get_available_voices_for_backend(backend: str | None = None) -> list[str]:
|
| 331 |
"""Return the curated voice list for a backend selector value."""
|
| 332 |
normalized_backend = get_backend_choice() if backend is None else _normalize_backend_provider(backend)
|
| 333 |
if normalized_backend == GEMINI_BACKEND:
|
| 334 |
return list(GEMINI_AVAILABLE_VOICES)
|
|
|
|
|
|
|
| 335 |
return list(AVAILABLE_VOICES)
|
| 336 |
|
| 337 |
|
|
@@ -341,6 +506,41 @@ def get_default_voice_for_backend(backend: str | None = None) -> str:
|
|
| 341 |
return DEFAULT_VOICE_BY_BACKEND[normalized_backend]
|
| 342 |
|
| 343 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 344 |
def is_gemini_model() -> bool:
|
| 345 |
"""Return True if the configured MODEL_NAME is a Gemini Live model."""
|
| 346 |
return get_backend_choice() == GEMINI_BACKEND
|
|
|
|
| 2 |
import sys
|
| 3 |
import logging
|
| 4 |
from pathlib import Path
|
| 5 |
+
from dataclasses import dataclass
|
| 6 |
+
from urllib.parse import urlsplit, parse_qsl, urlunsplit
|
| 7 |
from importlib.resources import files
|
| 8 |
|
| 9 |
from dotenv import find_dotenv, load_dotenv
|
|
|
|
| 58 |
"shimmer",
|
| 59 |
"verse",
|
| 60 |
]
|
| 61 |
+
OPENAI_DEFAULT_VOICE = "cedar"
|
| 62 |
+
|
| 63 |
+
# Qwen3-TTS CustomVoice speaker catalog from the deployed Hugging Face backend.
|
| 64 |
+
HF_AVAILABLE_VOICES: list[str] = [
|
| 65 |
+
"Aiden",
|
| 66 |
+
"Ryan",
|
| 67 |
+
"Dylan",
|
| 68 |
+
"Eric",
|
| 69 |
+
"Ono_Anna",
|
| 70 |
+
"Serena",
|
| 71 |
+
"Sohee",
|
| 72 |
+
"Uncle_Fu",
|
| 73 |
+
"Vivian",
|
| 74 |
+
]
|
| 75 |
|
| 76 |
# Voices supported by the Gemini Live API
|
| 77 |
GEMINI_AVAILABLE_VOICES: list[str] = [
|
|
|
|
| 87 |
|
| 88 |
OPENAI_BACKEND = "openai"
|
| 89 |
GEMINI_BACKEND = "gemini"
|
| 90 |
+
HF_BACKEND = "huggingface"
|
| 91 |
+
DEFAULT_BACKEND_PROVIDER = HF_BACKEND
|
| 92 |
+
HF_REALTIME_CONNECTION_MODE_ENV = "HF_REALTIME_CONNECTION_MODE"
|
| 93 |
+
HF_REALTIME_WS_URL_ENV = "HF_REALTIME_WS_URL"
|
| 94 |
+
HF_LOCAL_CONNECTION_MODE = "local"
|
| 95 |
+
HF_DEPLOYED_CONNECTION_MODE = "deployed"
|
| 96 |
+
HF_REALTIME_SESSION_PROXY_URL = "https://pollen-robotics-reachy-mini-realtime-url.hf.space/session"
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
@dataclass(frozen=True)
|
| 100 |
+
class HFBackendDefaults:
|
| 101 |
+
"""Defaults for the Hugging Face realtime backend."""
|
| 102 |
+
|
| 103 |
+
connection_mode: str = HF_DEPLOYED_CONNECTION_MODE
|
| 104 |
+
# App-managed Hugging Face Space proxy. The Space forwards to the current
|
| 105 |
+
# session allocator, so allocator changes do not require app releases.
|
| 106 |
+
# Users who need a custom target should use HF_REALTIME_CONNECTION_MODE=local
|
| 107 |
+
# with HF_REALTIME_WS_URL.
|
| 108 |
+
session_url: str = HF_REALTIME_SESSION_PROXY_URL
|
| 109 |
+
voice: str = "Aiden"
|
| 110 |
+
model_name: str = ""
|
| 111 |
+
direct_port: int = 8765
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
HF_DEFAULTS = HFBackendDefaults()
|
| 115 |
DEFAULT_MODEL_NAME_BY_BACKEND = {
|
| 116 |
OPENAI_BACKEND: "gpt-realtime",
|
| 117 |
GEMINI_BACKEND: "gemini-3.1-flash-live-preview",
|
| 118 |
+
HF_BACKEND: HF_DEFAULTS.model_name,
|
| 119 |
+
}
|
| 120 |
+
BACKEND_LABEL_BY_PROVIDER = {
|
| 121 |
+
OPENAI_BACKEND: "OpenAI Realtime",
|
| 122 |
+
GEMINI_BACKEND: "Gemini Live",
|
| 123 |
+
HF_BACKEND: "Hugging Face",
|
| 124 |
}
|
| 125 |
DEFAULT_VOICE_BY_BACKEND = {
|
| 126 |
+
OPENAI_BACKEND: OPENAI_DEFAULT_VOICE,
|
| 127 |
GEMINI_BACKEND: "Kore",
|
| 128 |
+
HF_BACKEND: HF_DEFAULTS.voice,
|
| 129 |
}
|
| 130 |
|
| 131 |
logger = logging.getLogger(__name__)
|
|
|
|
| 141 |
backend_provider: str | None = None,
|
| 142 |
model_name: str | None = None,
|
| 143 |
) -> str:
|
| 144 |
+
"""Normalize the configured backend provider."""
|
| 145 |
candidate = (backend_provider or "").strip().lower()
|
| 146 |
if candidate in DEFAULT_MODEL_NAME_BY_BACKEND:
|
| 147 |
return candidate
|
| 148 |
+
if candidate:
|
| 149 |
+
expected = ", ".join(sorted(DEFAULT_MODEL_NAME_BY_BACKEND))
|
| 150 |
+
raise ValueError(f"Invalid BACKEND_PROVIDER={backend_provider!r}. Expected one of: {expected}.")
|
| 151 |
return GEMINI_BACKEND if _is_gemini_model_name(model_name) else DEFAULT_BACKEND_PROVIDER
|
| 152 |
|
| 153 |
|
|
|
|
| 157 |
) -> str:
|
| 158 |
"""Return a model name that matches the selected backend provider."""
|
| 159 |
normalized_backend = _normalize_backend_provider(backend_provider, model_name)
|
| 160 |
+
if normalized_backend == HF_BACKEND:
|
| 161 |
+
return DEFAULT_MODEL_NAME_BY_BACKEND[HF_BACKEND]
|
| 162 |
+
|
| 163 |
candidate = (model_name or "").strip()
|
| 164 |
if candidate:
|
| 165 |
if normalized_backend == GEMINI_BACKEND and _is_gemini_model_name(candidate):
|
| 166 |
return candidate
|
| 167 |
+
if normalized_backend != GEMINI_BACKEND and not _is_gemini_model_name(candidate):
|
| 168 |
return candidate
|
| 169 |
logger.warning(
|
| 170 |
"MODEL_NAME=%r does not match BACKEND_PROVIDER=%r, using default %r",
|
|
|
|
| 195 |
return default
|
| 196 |
|
| 197 |
|
| 198 |
+
def _normalize_hf_connection_mode(value: str | None) -> str | None:
|
| 199 |
+
"""Normalize the Hugging Face connection mode, if explicitly configured."""
|
| 200 |
+
candidate = (value or "").strip().lower()
|
| 201 |
+
if not candidate:
|
| 202 |
+
return None
|
| 203 |
+
|
| 204 |
+
if candidate not in {HF_LOCAL_CONNECTION_MODE, HF_DEPLOYED_CONNECTION_MODE}:
|
| 205 |
+
logger.warning(
|
| 206 |
+
"Invalid %s=%r. Expected local or deployed.",
|
| 207 |
+
HF_REALTIME_CONNECTION_MODE_ENV,
|
| 208 |
+
value,
|
| 209 |
+
)
|
| 210 |
+
return None
|
| 211 |
+
return candidate
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
@dataclass(frozen=True)
|
| 215 |
+
class HFConnectionSelection:
|
| 216 |
+
"""Resolved Hugging Face connection mode and target availability."""
|
| 217 |
+
|
| 218 |
+
mode: str
|
| 219 |
+
has_target: bool
|
| 220 |
+
session_url: str | None = None
|
| 221 |
+
direct_ws_url: str | None = None
|
| 222 |
+
|
| 223 |
+
|
| 224 |
+
@dataclass(frozen=True)
|
| 225 |
+
class HFRealtimeURLParts:
|
| 226 |
+
"""Parsed Hugging Face realtime URL components used by UI and client setup."""
|
| 227 |
+
|
| 228 |
+
base_url: str
|
| 229 |
+
websocket_base_url: str
|
| 230 |
+
connect_query: dict[str, str]
|
| 231 |
+
host: str | None
|
| 232 |
+
port: int | None
|
| 233 |
+
has_realtime_path: bool
|
| 234 |
+
|
| 235 |
+
|
| 236 |
+
def parse_hf_realtime_url(realtime_url: str) -> HFRealtimeURLParts:
|
| 237 |
+
"""Parse a Hugging Face realtime URL into OpenAI-compatible client endpoints."""
|
| 238 |
+
parsed = urlsplit(realtime_url)
|
| 239 |
+
scheme = parsed.scheme.lower()
|
| 240 |
+
if scheme not in {"ws", "wss", "http", "https"}:
|
| 241 |
+
raise ValueError(
|
| 242 |
+
"Expected Hugging Face realtime URL to start with ws://, wss://, http://, or https://, "
|
| 243 |
+
f"got: {realtime_url}"
|
| 244 |
+
)
|
| 245 |
+
|
| 246 |
+
path = parsed.path.rstrip("/")
|
| 247 |
+
has_realtime_path = path.endswith("/realtime")
|
| 248 |
+
if has_realtime_path:
|
| 249 |
+
base_path = path[: -len("/realtime")]
|
| 250 |
+
else:
|
| 251 |
+
base_path = path
|
| 252 |
+
|
| 253 |
+
connect_query = {key: value for key, value in parse_qsl(parsed.query, keep_blank_values=True) if key != "model"}
|
| 254 |
+
http_scheme = "https" if scheme in {"wss", "https"} else "http"
|
| 255 |
+
websocket_scheme = "wss" if scheme in {"wss", "https"} else "ws"
|
| 256 |
+
base_url = urlunsplit((http_scheme, parsed.netloc, base_path, "", ""))
|
| 257 |
+
websocket_base_url = urlunsplit((websocket_scheme, parsed.netloc, base_path, "", ""))
|
| 258 |
+
return HFRealtimeURLParts(
|
| 259 |
+
base_url=base_url,
|
| 260 |
+
websocket_base_url=websocket_base_url,
|
| 261 |
+
connect_query=connect_query,
|
| 262 |
+
host=parsed.hostname,
|
| 263 |
+
port=parsed.port or HF_DEFAULTS.direct_port,
|
| 264 |
+
has_realtime_path=has_realtime_path,
|
| 265 |
+
)
|
| 266 |
+
|
| 267 |
+
|
| 268 |
+
def parse_hf_direct_target(ws_url: str | None) -> tuple[str | None, int | None]:
|
| 269 |
+
"""Extract host and port from a direct Hugging Face realtime URL."""
|
| 270 |
+
if not ws_url:
|
| 271 |
+
return None, None
|
| 272 |
+
try:
|
| 273 |
+
parsed = parse_hf_realtime_url(ws_url)
|
| 274 |
+
return parsed.host, parsed.port
|
| 275 |
+
except Exception:
|
| 276 |
+
return None, None
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
def build_hf_direct_ws_url(host: str, port: int) -> str:
|
| 280 |
+
"""Build the direct Hugging Face realtime websocket URL used by the app."""
|
| 281 |
+
return f"ws://{host}:{port}/v1/realtime"
|
| 282 |
+
|
| 283 |
+
|
| 284 |
def _collect_profile_names(profiles_root: Path) -> set[str]:
|
| 285 |
"""Return profile folder names from a profiles root directory."""
|
| 286 |
if not profiles_root.exists() or not profiles_root.is_dir():
|
|
|
|
| 357 |
os.getenv("MODEL_NAME"),
|
| 358 |
)
|
| 359 |
MODEL_NAME = _resolve_model_name(BACKEND_PROVIDER, os.getenv("MODEL_NAME"))
|
| 360 |
+
HF_REALTIME_CONNECTION_MODE = (
|
| 361 |
+
_normalize_hf_connection_mode(os.getenv(HF_REALTIME_CONNECTION_MODE_ENV)) or HF_DEFAULTS.connection_mode
|
| 362 |
+
)
|
| 363 |
+
# Deliberately ignore HF_REALTIME_SESSION_URL from the environment; the app-managed proxy is HF_DEFAULTS.session_url.
|
| 364 |
+
HF_REALTIME_SESSION_URL = HF_DEFAULTS.session_url
|
| 365 |
+
HF_REALTIME_WS_URL = os.getenv(HF_REALTIME_WS_URL_ENV)
|
| 366 |
HF_HOME = os.getenv("HF_HOME", "./cache")
|
| 367 |
LOCAL_VISION_MODEL = os.getenv("LOCAL_VISION_MODEL", "HuggingFaceTB/SmolVLM2-2.2B-Instruct")
|
| 368 |
HF_TOKEN = os.getenv("HF_TOKEN") # Optional, falls back to hf auth login if not set
|
| 369 |
|
| 370 |
logger.debug(
|
| 371 |
+
"Backend provider: %s, Model: %s, HF mode: %s, HF session URL set: %s, HF direct URL set: %s, HF_HOME: %s, Vision Model: %s",
|
| 372 |
BACKEND_PROVIDER,
|
| 373 |
MODEL_NAME,
|
| 374 |
+
HF_REALTIME_CONNECTION_MODE,
|
| 375 |
+
bool(HF_REALTIME_SESSION_URL and HF_REALTIME_SESSION_URL.strip()),
|
| 376 |
+
bool(HF_REALTIME_WS_URL and HF_REALTIME_WS_URL.strip()),
|
| 377 |
HF_HOME,
|
| 378 |
LOCAL_VISION_MODEL,
|
| 379 |
)
|
|
|
|
| 460 |
os.getenv("MODEL_NAME"),
|
| 461 |
)
|
| 462 |
config.MODEL_NAME = _resolve_model_name(config.BACKEND_PROVIDER, os.getenv("MODEL_NAME"))
|
| 463 |
+
config.HF_REALTIME_CONNECTION_MODE = (
|
| 464 |
+
_normalize_hf_connection_mode(os.getenv(HF_REALTIME_CONNECTION_MODE_ENV)) or HF_DEFAULTS.connection_mode
|
| 465 |
+
)
|
| 466 |
+
# Deliberately ignore HF_REALTIME_SESSION_URL from the environment; the app-managed proxy is HF_DEFAULTS.session_url.
|
| 467 |
+
config.HF_REALTIME_SESSION_URL = HF_DEFAULTS.session_url
|
| 468 |
+
config.HF_REALTIME_WS_URL = os.getenv(HF_REALTIME_WS_URL_ENV)
|
| 469 |
+
config.HF_HOME = os.getenv("HF_HOME", "./cache")
|
| 470 |
+
config.LOCAL_VISION_MODEL = os.getenv("LOCAL_VISION_MODEL", "HuggingFaceTB/SmolVLM2-2.2B-Instruct")
|
| 471 |
+
config.HF_TOKEN = os.getenv("HF_TOKEN")
|
| 472 |
config.REACHY_MINI_CUSTOM_PROFILE = LOCKED_PROFILE or os.getenv("REACHY_MINI_CUSTOM_PROFILE")
|
| 473 |
|
| 474 |
|
|
|
|
| 484 |
return DEFAULT_MODEL_NAME_BY_BACKEND[_normalize_backend_provider(backend)]
|
| 485 |
|
| 486 |
|
| 487 |
+
def get_backend_label(backend: str | None = None) -> str:
|
| 488 |
+
"""Return a human-readable label for a backend selector value."""
|
| 489 |
+
normalized_backend = get_backend_choice() if backend is None else _normalize_backend_provider(backend)
|
| 490 |
+
return BACKEND_LABEL_BY_PROVIDER[normalized_backend]
|
| 491 |
+
|
| 492 |
+
|
| 493 |
def get_available_voices_for_backend(backend: str | None = None) -> list[str]:
|
| 494 |
"""Return the curated voice list for a backend selector value."""
|
| 495 |
normalized_backend = get_backend_choice() if backend is None else _normalize_backend_provider(backend)
|
| 496 |
if normalized_backend == GEMINI_BACKEND:
|
| 497 |
return list(GEMINI_AVAILABLE_VOICES)
|
| 498 |
+
if normalized_backend == HF_BACKEND:
|
| 499 |
+
return list(HF_AVAILABLE_VOICES)
|
| 500 |
return list(AVAILABLE_VOICES)
|
| 501 |
|
| 502 |
|
|
|
|
| 506 |
return DEFAULT_VOICE_BY_BACKEND[normalized_backend]
|
| 507 |
|
| 508 |
|
| 509 |
+
def get_hf_session_url() -> str | None:
|
| 510 |
+
"""Return the built-in Hugging Face session proxy URL, if any."""
|
| 511 |
+
value = (getattr(config, "HF_REALTIME_SESSION_URL", None) or "").strip()
|
| 512 |
+
return value or None
|
| 513 |
+
|
| 514 |
+
|
| 515 |
+
def get_hf_direct_ws_url() -> str | None:
|
| 516 |
+
"""Return the configured direct Hugging Face realtime URL, if any."""
|
| 517 |
+
value = (getattr(config, "HF_REALTIME_WS_URL", None) or "").strip()
|
| 518 |
+
return value or None
|
| 519 |
+
|
| 520 |
+
|
| 521 |
+
def get_hf_connection_selection() -> HFConnectionSelection:
|
| 522 |
+
"""Resolve the selected Hugging Face connection mode and whether it is usable."""
|
| 523 |
+
session_url = get_hf_session_url()
|
| 524 |
+
direct_ws_url = get_hf_direct_ws_url()
|
| 525 |
+
mode = _normalize_hf_connection_mode(getattr(config, "HF_REALTIME_CONNECTION_MODE", None))
|
| 526 |
+
if mode is None:
|
| 527 |
+
raise RuntimeError(f"{HF_REALTIME_CONNECTION_MODE_ENV} must be set to local or deployed.")
|
| 528 |
+
|
| 529 |
+
target = direct_ws_url if mode == HF_LOCAL_CONNECTION_MODE else session_url
|
| 530 |
+
|
| 531 |
+
return HFConnectionSelection(
|
| 532 |
+
mode=mode,
|
| 533 |
+
has_target=bool(target),
|
| 534 |
+
session_url=session_url,
|
| 535 |
+
direct_ws_url=direct_ws_url,
|
| 536 |
+
)
|
| 537 |
+
|
| 538 |
+
|
| 539 |
+
def has_hf_realtime_target() -> bool:
|
| 540 |
+
"""Return whether Hugging Face has a target for the selected mode."""
|
| 541 |
+
return get_hf_connection_selection().has_target
|
| 542 |
+
|
| 543 |
+
|
| 544 |
def is_gemini_model() -> bool:
|
| 545 |
"""Return True if the configured MODEL_NAME is a Gemini Live model."""
|
| 546 |
return get_backend_choice() == GEMINI_BACKEND
|
src/reachy_mini_conversation_app/console.py
CHANGED
|
@@ -24,24 +24,31 @@ from scipy.signal import resample
|
|
| 24 |
from reachy_mini import ReachyMini
|
| 25 |
from reachy_mini.media.media_manager import MediaBackend
|
| 26 |
from reachy_mini_conversation_app.config import (
|
|
|
|
| 27 |
GEMINI_BACKEND,
|
| 28 |
LOCKED_PROFILE,
|
| 29 |
OPENAI_BACKEND,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
config,
|
| 31 |
get_backend_choice,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 32 |
get_model_name_for_backend,
|
|
|
|
| 33 |
refresh_runtime_config_from_env,
|
| 34 |
)
|
| 35 |
-
from reachy_mini_conversation_app.
|
|
|
|
|
|
|
| 36 |
from reachy_mini_conversation_app.headless_personality_ui import mount_personality_routes
|
| 37 |
|
| 38 |
|
| 39 |
-
try:
|
| 40 |
-
from reachy_mini_conversation_app.gemini_live import GeminiLiveHandler
|
| 41 |
-
except ImportError:
|
| 42 |
-
GeminiLiveHandler = None # type: ignore[misc,assignment]
|
| 43 |
-
|
| 44 |
-
|
| 45 |
try:
|
| 46 |
# FastAPI is provided by the Reachy Mini Apps runtime
|
| 47 |
from fastapi import FastAPI, Response
|
|
@@ -58,6 +65,17 @@ except Exception: # pragma: no cover - only loaded when settings_app is used
|
|
| 58 |
|
| 59 |
logger = logging.getLogger(__name__)
|
| 60 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 61 |
|
| 62 |
def _estimate_pending_playback_seconds(robot: ReachyMini) -> float:
|
| 63 |
"""Best-effort estimate of audio still queued in the local player."""
|
|
@@ -84,13 +102,13 @@ class LocalStream:
|
|
| 84 |
|
| 85 |
def __init__(
|
| 86 |
self,
|
| 87 |
-
handler:
|
| 88 |
robot: ReachyMini,
|
| 89 |
*,
|
| 90 |
settings_app: Optional[FastAPI] = None,
|
| 91 |
instance_path: Optional[str] = None,
|
| 92 |
):
|
| 93 |
-
"""Initialize the stream with
|
| 94 |
|
| 95 |
- ``settings_app``: the Reachy Mini Apps FastAPI to attach settings endpoints.
|
| 96 |
- ``instance_path``: directory where per-instance ``.env`` should be stored.
|
|
@@ -105,6 +123,7 @@ class LocalStream:
|
|
| 105 |
self._instance_path: Optional[str] = instance_path
|
| 106 |
self._settings_initialized = False
|
| 107 |
self._asyncio_loop = None
|
|
|
|
| 108 |
|
| 109 |
# ---- Settings UI ----
|
| 110 |
def _read_env_lines(self, env_path: Path) -> list[str]:
|
|
@@ -143,8 +162,7 @@ class LocalStream:
|
|
| 143 |
|
| 144 |
def _active_backend(self) -> str:
|
| 145 |
"""Return the backend family of the currently running handler."""
|
| 146 |
-
|
| 147 |
-
return GEMINI_BACKEND if "gemini" in handler_name else OPENAI_BACKEND
|
| 148 |
|
| 149 |
@staticmethod
|
| 150 |
def _has_key(value: Optional[str]) -> bool:
|
|
@@ -155,8 +173,19 @@ class LocalStream:
|
|
| 155 |
"""Return whether the requested backend has its required credential."""
|
| 156 |
if backend == GEMINI_BACKEND:
|
| 157 |
return self._has_key(config.GEMINI_API_KEY)
|
|
|
|
|
|
|
| 158 |
return self._has_key(config.OPENAI_API_KEY)
|
| 159 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 160 |
def _persist_env_value(self, env_name: str, value: str) -> None:
|
| 161 |
"""Persist a non-empty environment value in memory and in the instance `.env`."""
|
| 162 |
self._persist_env_values({env_name: value})
|
|
@@ -204,6 +233,48 @@ class LocalStream:
|
|
| 204 |
except Exception as e:
|
| 205 |
logger.warning("Failed to persist %s: %s", ", ".join(sorted(normalized_updates)), e)
|
| 206 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 207 |
def _persist_api_key(self, key: str) -> None:
|
| 208 |
"""Persist OPENAI_API_KEY to environment and instance `.env`."""
|
| 209 |
self._persist_env_value("OPENAI_API_KEY", key)
|
|
@@ -217,17 +288,28 @@ class LocalStream:
|
|
| 217 |
current_backend = get_backend_choice()
|
| 218 |
current_model_name = (os.getenv("MODEL_NAME") or "").strip()
|
| 219 |
updates = {"BACKEND_PROVIDER": backend}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 220 |
if current_model_name and current_model_name != get_model_name_for_backend(current_backend):
|
| 221 |
updates["MODEL_NAME"] = current_model_name
|
| 222 |
else:
|
| 223 |
updates["MODEL_NAME"] = get_model_name_for_backend(backend)
|
| 224 |
self._persist_env_values(updates)
|
| 225 |
|
| 226 |
-
def _persist_personality(self, profile: Optional[str]) -> None:
|
| 227 |
-
"""Persist
|
| 228 |
if LOCKED_PROFILE is not None:
|
| 229 |
return
|
| 230 |
selection = (profile or "").strip() or None
|
|
|
|
| 231 |
try:
|
| 232 |
from reachy_mini_conversation_app.config import set_custom_profile
|
| 233 |
|
|
@@ -238,48 +320,19 @@ class LocalStream:
|
|
| 238 |
if not self._instance_path:
|
| 239 |
return
|
| 240 |
try:
|
| 241 |
-
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
-
|
| 248 |
-
else:
|
| 249 |
-
lines.pop(i)
|
| 250 |
-
replaced = True
|
| 251 |
-
break
|
| 252 |
-
if selection and not replaced:
|
| 253 |
-
lines.append(f"REACHY_MINI_CUSTOM_PROFILE={selection}")
|
| 254 |
-
if selection is None and not env_path.exists():
|
| 255 |
-
return
|
| 256 |
-
final_text = "\n".join(lines) + "\n"
|
| 257 |
-
env_path.write_text(final_text, encoding="utf-8")
|
| 258 |
-
logger.info("Persisted startup personality to %s", env_path)
|
| 259 |
-
try:
|
| 260 |
-
from dotenv import load_dotenv
|
| 261 |
-
|
| 262 |
-
load_dotenv(dotenv_path=str(env_path), override=True)
|
| 263 |
-
except Exception:
|
| 264 |
-
pass
|
| 265 |
except Exception as e:
|
| 266 |
-
logger.warning("Failed to persist
|
| 267 |
|
| 268 |
def _read_persisted_personality(self) -> Optional[str]:
|
| 269 |
-
"""Read
|
| 270 |
-
|
| 271 |
-
return None
|
| 272 |
-
env_path = Path(self._instance_path) / ".env"
|
| 273 |
-
try:
|
| 274 |
-
if env_path.exists():
|
| 275 |
-
for ln in env_path.read_text(encoding="utf-8").splitlines():
|
| 276 |
-
if ln.strip().startswith("REACHY_MINI_CUSTOM_PROFILE="):
|
| 277 |
-
_, _, val = ln.partition("=")
|
| 278 |
-
v = val.strip()
|
| 279 |
-
return v or None
|
| 280 |
-
except Exception:
|
| 281 |
-
pass
|
| 282 |
-
return None
|
| 283 |
|
| 284 |
def _init_settings_ui_if_needed(self) -> None:
|
| 285 |
"""Attach minimal settings UI to the settings app.
|
|
@@ -308,14 +361,26 @@ class LocalStream:
|
|
| 308 |
class BackendPayload(BaseModel):
|
| 309 |
backend: str
|
| 310 |
api_key: Optional[str] = None
|
|
|
|
|
|
|
|
|
|
| 311 |
|
| 312 |
def _status_payload() -> dict[str, object]:
|
| 313 |
backend_provider = get_backend_choice()
|
| 314 |
active_backend = self._active_backend()
|
| 315 |
has_openai_key = self._has_required_key(OPENAI_BACKEND)
|
| 316 |
has_gemini_key = self._has_required_key(GEMINI_BACKEND)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 317 |
can_proceed_with_openai = has_openai_key
|
| 318 |
can_proceed_with_gemini = has_gemini_key
|
|
|
|
| 319 |
can_proceed = self._has_required_key(active_backend)
|
| 320 |
requires_restart = backend_provider != active_backend
|
| 321 |
return {
|
|
@@ -324,9 +389,16 @@ class LocalStream:
|
|
| 324 |
"has_key": can_proceed,
|
| 325 |
"has_openai_key": has_openai_key,
|
| 326 |
"has_gemini_key": has_gemini_key,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 327 |
"can_proceed": can_proceed,
|
| 328 |
"can_proceed_with_openai": can_proceed_with_openai,
|
| 329 |
"can_proceed_with_gemini": can_proceed_with_gemini,
|
|
|
|
| 330 |
"requires_restart": requires_restart,
|
| 331 |
}
|
| 332 |
|
|
@@ -367,7 +439,7 @@ class LocalStream:
|
|
| 367 |
@self._settings_app.post("/backend_config")
|
| 368 |
def _set_backend(payload: BackendPayload) -> JSONResponse:
|
| 369 |
backend = payload.backend.strip().lower()
|
| 370 |
-
if backend not in {OPENAI_BACKEND, GEMINI_BACKEND}:
|
| 371 |
return JSONResponse({"ok": False, "error": "invalid_backend"}, status_code=400)
|
| 372 |
|
| 373 |
api_key = (payload.api_key or "").strip()
|
|
@@ -378,6 +450,28 @@ class LocalStream:
|
|
| 378 |
self._persist_api_key(api_key)
|
| 379 |
if backend == GEMINI_BACKEND and api_key:
|
| 380 |
self._persist_gemini_api_key(api_key)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 381 |
|
| 382 |
self._persist_backend_choice(backend)
|
| 383 |
payload_data = _status_payload()
|
|
@@ -443,29 +537,20 @@ class LocalStream:
|
|
| 443 |
|
| 444 |
active_backend = self._active_backend()
|
| 445 |
|
| 446 |
-
# If key is still missing, try to download one from HuggingFace (OpenAI only)
|
| 447 |
-
if active_backend == OPENAI_BACKEND and not self._has_required_key(active_backend):
|
| 448 |
-
logger.info("OPENAI_API_KEY not set, attempting to download from HuggingFace...")
|
| 449 |
-
try:
|
| 450 |
-
from gradio_client import Client
|
| 451 |
-
|
| 452 |
-
client = Client("HuggingFaceM4/gradium_setup", verbose=False)
|
| 453 |
-
key, _ = client.predict(api_name="/claim_b_key")
|
| 454 |
-
if key and key.strip():
|
| 455 |
-
logger.info("Successfully downloaded API key from HuggingFace")
|
| 456 |
-
# Persist it immediately
|
| 457 |
-
self._persist_api_key(key)
|
| 458 |
-
except Exception as e:
|
| 459 |
-
logger.warning(f"Failed to download API key from HuggingFace: {e}")
|
| 460 |
-
|
| 461 |
# Always expose settings UI if a settings app is available
|
| 462 |
-
# (do this AFTER loading
|
| 463 |
self._init_settings_ui_if_needed()
|
| 464 |
|
| 465 |
# If key is still missing -> wait until provided via the settings UI
|
| 466 |
if not self._has_required_key(active_backend):
|
| 467 |
-
|
| 468 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 469 |
# Poll until the key becomes available (set via the settings UI)
|
| 470 |
try:
|
| 471 |
while not self._has_required_key(active_backend):
|
|
@@ -478,6 +563,7 @@ class LocalStream:
|
|
| 478 |
self._robot.media.start_recording()
|
| 479 |
self._robot.media.start_playing()
|
| 480 |
time.sleep(1) # give some time to the pipelines to start
|
|
|
|
| 481 |
|
| 482 |
async def runner() -> None:
|
| 483 |
# Capture loop for cross-thread personality actions
|
|
@@ -546,7 +632,12 @@ class LocalStream:
|
|
| 546 |
backend = getattr(self._robot.media, "backend", None)
|
| 547 |
audio = getattr(self._robot.media, "audio", None)
|
| 548 |
if audio is not None:
|
| 549 |
-
if
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 550 |
audio.clear_player()
|
| 551 |
elif (
|
| 552 |
backend == MediaBackend.WEBRTC
|
|
|
|
| 24 |
from reachy_mini import ReachyMini
|
| 25 |
from reachy_mini.media.media_manager import MediaBackend
|
| 26 |
from reachy_mini_conversation_app.config import (
|
| 27 |
+
HF_BACKEND,
|
| 28 |
GEMINI_BACKEND,
|
| 29 |
LOCKED_PROFILE,
|
| 30 |
OPENAI_BACKEND,
|
| 31 |
+
HF_REALTIME_WS_URL_ENV,
|
| 32 |
+
HF_LOCAL_CONNECTION_MODE,
|
| 33 |
+
HF_DEPLOYED_CONNECTION_MODE,
|
| 34 |
+
HF_REALTIME_CONNECTION_MODE_ENV,
|
| 35 |
config,
|
| 36 |
get_backend_choice,
|
| 37 |
+
get_hf_session_url,
|
| 38 |
+
get_hf_direct_ws_url,
|
| 39 |
+
build_hf_direct_ws_url,
|
| 40 |
+
has_hf_realtime_target,
|
| 41 |
+
parse_hf_direct_target,
|
| 42 |
get_model_name_for_backend,
|
| 43 |
+
get_hf_connection_selection,
|
| 44 |
refresh_runtime_config_from_env,
|
| 45 |
)
|
| 46 |
+
from reachy_mini_conversation_app.startup_settings import read_startup_settings, write_startup_settings
|
| 47 |
+
from reachy_mini_conversation_app.audio.startup_config import apply_audio_startup_config
|
| 48 |
+
from reachy_mini_conversation_app.conversation_handler import ConversationHandler
|
| 49 |
from reachy_mini_conversation_app.headless_personality_ui import mount_personality_routes
|
| 50 |
|
| 51 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 52 |
try:
|
| 53 |
# FastAPI is provided by the Reachy Mini Apps runtime
|
| 54 |
from fastapi import FastAPI, Response
|
|
|
|
| 65 |
|
| 66 |
logger = logging.getLogger(__name__)
|
| 67 |
|
| 68 |
+
LOCAL_PLAYER_BACKEND = (
|
| 69 |
+
getattr(MediaBackend, "LOCAL", None)
|
| 70 |
+
or getattr(MediaBackend, "GSTREAMER", None)
|
| 71 |
+
or getattr(MediaBackend, "DEFAULT", None)
|
| 72 |
+
)
|
| 73 |
+
|
| 74 |
+
LEGACY_STARTUP_ENV_NAMES = (
|
| 75 |
+
"REACHY_MINI_CUSTOM_PROFILE",
|
| 76 |
+
"REACHY_MINI_VOICE_OVERRIDE",
|
| 77 |
+
)
|
| 78 |
+
|
| 79 |
|
| 80 |
def _estimate_pending_playback_seconds(robot: ReachyMini) -> float:
|
| 81 |
"""Best-effort estimate of audio still queued in the local player."""
|
|
|
|
| 102 |
|
| 103 |
def __init__(
|
| 104 |
self,
|
| 105 |
+
handler: ConversationHandler,
|
| 106 |
robot: ReachyMini,
|
| 107 |
*,
|
| 108 |
settings_app: Optional[FastAPI] = None,
|
| 109 |
instance_path: Optional[str] = None,
|
| 110 |
):
|
| 111 |
+
"""Initialize the stream with a realtime handler and pipelines.
|
| 112 |
|
| 113 |
- ``settings_app``: the Reachy Mini Apps FastAPI to attach settings endpoints.
|
| 114 |
- ``instance_path``: directory where per-instance ``.env`` should be stored.
|
|
|
|
| 123 |
self._instance_path: Optional[str] = instance_path
|
| 124 |
self._settings_initialized = False
|
| 125 |
self._asyncio_loop = None
|
| 126 |
+
self._active_backend_name = get_backend_choice()
|
| 127 |
|
| 128 |
# ---- Settings UI ----
|
| 129 |
def _read_env_lines(self, env_path: Path) -> list[str]:
|
|
|
|
| 162 |
|
| 163 |
def _active_backend(self) -> str:
|
| 164 |
"""Return the backend family of the currently running handler."""
|
| 165 |
+
return self._active_backend_name
|
|
|
|
| 166 |
|
| 167 |
@staticmethod
|
| 168 |
def _has_key(value: Optional[str]) -> bool:
|
|
|
|
| 173 |
"""Return whether the requested backend has its required credential."""
|
| 174 |
if backend == GEMINI_BACKEND:
|
| 175 |
return self._has_key(config.GEMINI_API_KEY)
|
| 176 |
+
if backend == HF_BACKEND:
|
| 177 |
+
return has_hf_realtime_target()
|
| 178 |
return self._has_key(config.OPENAI_API_KEY)
|
| 179 |
|
| 180 |
+
@staticmethod
|
| 181 |
+
def _requirement_name(backend: str) -> str:
|
| 182 |
+
"""Return the env var users need for a backend, if any."""
|
| 183 |
+
if backend == GEMINI_BACKEND:
|
| 184 |
+
return "GEMINI_API_KEY"
|
| 185 |
+
if backend == HF_BACKEND:
|
| 186 |
+
return HF_REALTIME_WS_URL_ENV
|
| 187 |
+
return "OPENAI_API_KEY"
|
| 188 |
+
|
| 189 |
def _persist_env_value(self, env_name: str, value: str) -> None:
|
| 190 |
"""Persist a non-empty environment value in memory and in the instance `.env`."""
|
| 191 |
self._persist_env_values({env_name: value})
|
|
|
|
| 233 |
except Exception as e:
|
| 234 |
logger.warning("Failed to persist %s: %s", ", ".join(sorted(normalized_updates)), e)
|
| 235 |
|
| 236 |
+
def _remove_persisted_env_values(self, env_names: tuple[str, ...]) -> None:
|
| 237 |
+
"""Remove keys from the instance `.env` without mutating the current runtime."""
|
| 238 |
+
normalized_names = tuple(sorted({name.strip() for name in env_names if name and name.strip()}))
|
| 239 |
+
if not normalized_names or not self._instance_path:
|
| 240 |
+
return
|
| 241 |
+
|
| 242 |
+
env_path = Path(self._instance_path) / ".env"
|
| 243 |
+
if not env_path.exists():
|
| 244 |
+
return
|
| 245 |
+
|
| 246 |
+
try:
|
| 247 |
+
lines = env_path.read_text(encoding="utf-8").splitlines()
|
| 248 |
+
filtered_lines = [
|
| 249 |
+
line
|
| 250 |
+
for line in lines
|
| 251 |
+
if not any(line.strip().startswith(f"{env_name}=") for env_name in normalized_names)
|
| 252 |
+
]
|
| 253 |
+
if filtered_lines == lines:
|
| 254 |
+
return
|
| 255 |
+
|
| 256 |
+
final_text = "\n".join(filtered_lines)
|
| 257 |
+
if final_text:
|
| 258 |
+
final_text += "\n"
|
| 259 |
+
env_path.write_text(final_text, encoding="utf-8")
|
| 260 |
+
logger.info("Removed %s from %s", ", ".join(normalized_names), env_path)
|
| 261 |
+
except Exception as e:
|
| 262 |
+
logger.warning("Failed to remove %s: %s", ", ".join(normalized_names), e)
|
| 263 |
+
|
| 264 |
+
def _persist_hf_direct_connection(self, host: str, port: int) -> None:
|
| 265 |
+
"""Persist a direct Hugging Face websocket target."""
|
| 266 |
+
self._persist_env_values(
|
| 267 |
+
{
|
| 268 |
+
HF_REALTIME_CONNECTION_MODE_ENV: HF_LOCAL_CONNECTION_MODE,
|
| 269 |
+
HF_REALTIME_WS_URL_ENV: build_hf_direct_ws_url(host, port),
|
| 270 |
+
}
|
| 271 |
+
)
|
| 272 |
+
|
| 273 |
+
def _persist_hf_allocator_connection(self) -> None:
|
| 274 |
+
"""Persist the deployed Hugging Face allocator mode."""
|
| 275 |
+
self._persist_env_value(HF_REALTIME_CONNECTION_MODE_ENV, HF_DEPLOYED_CONNECTION_MODE)
|
| 276 |
+
self._remove_persisted_env_values(("HF_REALTIME_SESSION_URL",))
|
| 277 |
+
|
| 278 |
def _persist_api_key(self, key: str) -> None:
|
| 279 |
"""Persist OPENAI_API_KEY to environment and instance `.env`."""
|
| 280 |
self._persist_env_value("OPENAI_API_KEY", key)
|
|
|
|
| 288 |
current_backend = get_backend_choice()
|
| 289 |
current_model_name = (os.getenv("MODEL_NAME") or "").strip()
|
| 290 |
updates = {"BACKEND_PROVIDER": backend}
|
| 291 |
+
if backend == HF_BACKEND:
|
| 292 |
+
self._persist_env_values(updates)
|
| 293 |
+
try:
|
| 294 |
+
os.environ.pop("MODEL_NAME", None)
|
| 295 |
+
except Exception:
|
| 296 |
+
pass
|
| 297 |
+
self._remove_persisted_env_values(("MODEL_NAME",))
|
| 298 |
+
refresh_runtime_config_from_env()
|
| 299 |
+
return
|
| 300 |
+
|
| 301 |
if current_model_name and current_model_name != get_model_name_for_backend(current_backend):
|
| 302 |
updates["MODEL_NAME"] = current_model_name
|
| 303 |
else:
|
| 304 |
updates["MODEL_NAME"] = get_model_name_for_backend(backend)
|
| 305 |
self._persist_env_values(updates)
|
| 306 |
|
| 307 |
+
def _persist_personality(self, profile: Optional[str], voice_override: Optional[str] = None) -> None:
|
| 308 |
+
"""Persist startup profile and voice in instance-local UI settings."""
|
| 309 |
if LOCKED_PROFILE is not None:
|
| 310 |
return
|
| 311 |
selection = (profile or "").strip() or None
|
| 312 |
+
normalized_voice_override = (voice_override or "").strip() or None
|
| 313 |
try:
|
| 314 |
from reachy_mini_conversation_app.config import set_custom_profile
|
| 315 |
|
|
|
|
| 320 |
if not self._instance_path:
|
| 321 |
return
|
| 322 |
try:
|
| 323 |
+
write_startup_settings(
|
| 324 |
+
self._instance_path,
|
| 325 |
+
profile=selection,
|
| 326 |
+
voice=normalized_voice_override,
|
| 327 |
+
)
|
| 328 |
+
self._remove_persisted_env_values(LEGACY_STARTUP_ENV_NAMES)
|
| 329 |
+
logger.info("Persisted startup personality settings to %s", Path(self._instance_path))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 330 |
except Exception as e:
|
| 331 |
+
logger.warning("Failed to persist startup personality settings: %s", e)
|
| 332 |
|
| 333 |
def _read_persisted_personality(self) -> Optional[str]:
|
| 334 |
+
"""Read the saved startup personality from instance-local UI settings."""
|
| 335 |
+
return read_startup_settings(self._instance_path).profile
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 336 |
|
| 337 |
def _init_settings_ui_if_needed(self) -> None:
|
| 338 |
"""Attach minimal settings UI to the settings app.
|
|
|
|
| 361 |
class BackendPayload(BaseModel):
|
| 362 |
backend: str
|
| 363 |
api_key: Optional[str] = None
|
| 364 |
+
hf_mode: Optional[str] = None
|
| 365 |
+
hf_host: Optional[str] = None
|
| 366 |
+
hf_port: Optional[int] = None
|
| 367 |
|
| 368 |
def _status_payload() -> dict[str, object]:
|
| 369 |
backend_provider = get_backend_choice()
|
| 370 |
active_backend = self._active_backend()
|
| 371 |
has_openai_key = self._has_required_key(OPENAI_BACKEND)
|
| 372 |
has_gemini_key = self._has_required_key(GEMINI_BACKEND)
|
| 373 |
+
hf_session_url = get_hf_session_url()
|
| 374 |
+
hf_ws_url = get_hf_direct_ws_url()
|
| 375 |
+
hf_direct_host, hf_direct_port = parse_hf_direct_target(hf_ws_url)
|
| 376 |
+
has_hf_session_url = bool(hf_session_url)
|
| 377 |
+
has_hf_ws_url = bool(hf_ws_url)
|
| 378 |
+
hf_connection_selection = get_hf_connection_selection()
|
| 379 |
+
hf_connection_mode = hf_connection_selection.mode
|
| 380 |
+
has_hf_connection = hf_connection_selection.has_target
|
| 381 |
can_proceed_with_openai = has_openai_key
|
| 382 |
can_proceed_with_gemini = has_gemini_key
|
| 383 |
+
can_proceed_with_hf = has_hf_connection
|
| 384 |
can_proceed = self._has_required_key(active_backend)
|
| 385 |
requires_restart = backend_provider != active_backend
|
| 386 |
return {
|
|
|
|
| 389 |
"has_key": can_proceed,
|
| 390 |
"has_openai_key": has_openai_key,
|
| 391 |
"has_gemini_key": has_gemini_key,
|
| 392 |
+
"has_hf_session_url": has_hf_session_url,
|
| 393 |
+
"has_hf_ws_url": has_hf_ws_url,
|
| 394 |
+
"has_hf_connection": has_hf_connection,
|
| 395 |
+
"hf_connection_mode": hf_connection_mode,
|
| 396 |
+
"hf_direct_host": hf_direct_host,
|
| 397 |
+
"hf_direct_port": hf_direct_port,
|
| 398 |
"can_proceed": can_proceed,
|
| 399 |
"can_proceed_with_openai": can_proceed_with_openai,
|
| 400 |
"can_proceed_with_gemini": can_proceed_with_gemini,
|
| 401 |
+
"can_proceed_with_hf": can_proceed_with_hf,
|
| 402 |
"requires_restart": requires_restart,
|
| 403 |
}
|
| 404 |
|
|
|
|
| 439 |
@self._settings_app.post("/backend_config")
|
| 440 |
def _set_backend(payload: BackendPayload) -> JSONResponse:
|
| 441 |
backend = payload.backend.strip().lower()
|
| 442 |
+
if backend not in {OPENAI_BACKEND, GEMINI_BACKEND, HF_BACKEND}:
|
| 443 |
return JSONResponse({"ok": False, "error": "invalid_backend"}, status_code=400)
|
| 444 |
|
| 445 |
api_key = (payload.api_key or "").strip()
|
|
|
|
| 450 |
self._persist_api_key(api_key)
|
| 451 |
if backend == GEMINI_BACKEND and api_key:
|
| 452 |
self._persist_gemini_api_key(api_key)
|
| 453 |
+
if backend == HF_BACKEND:
|
| 454 |
+
hf_selection = get_hf_connection_selection()
|
| 455 |
+
hf_mode = (payload.hf_mode or hf_selection.mode).strip().lower()
|
| 456 |
+
if hf_mode == HF_LOCAL_CONNECTION_MODE:
|
| 457 |
+
existing_host, existing_port = parse_hf_direct_target(hf_selection.direct_ws_url)
|
| 458 |
+
host = (payload.hf_host or "").strip() or existing_host or ""
|
| 459 |
+
if not host:
|
| 460 |
+
return JSONResponse({"ok": False, "error": "empty_hf_host"}, status_code=400)
|
| 461 |
+
if "://" in host or "/" in host or "?" in host or "#" in host:
|
| 462 |
+
return JSONResponse({"ok": False, "error": "invalid_hf_host"}, status_code=400)
|
| 463 |
+
|
| 464 |
+
port = payload.hf_port if payload.hf_port is not None else existing_port or 8765
|
| 465 |
+
if port < 1 or port > 65535:
|
| 466 |
+
return JSONResponse({"ok": False, "error": "invalid_hf_port"}, status_code=400)
|
| 467 |
+
|
| 468 |
+
self._persist_hf_direct_connection(host, port)
|
| 469 |
+
elif hf_mode == HF_DEPLOYED_CONNECTION_MODE:
|
| 470 |
+
if not bool(get_hf_session_url()):
|
| 471 |
+
return JSONResponse({"ok": False, "error": "missing_hf_session_url"}, status_code=400)
|
| 472 |
+
self._persist_hf_allocator_connection()
|
| 473 |
+
else:
|
| 474 |
+
return JSONResponse({"ok": False, "error": "invalid_hf_mode"}, status_code=400)
|
| 475 |
|
| 476 |
self._persist_backend_choice(backend)
|
| 477 |
payload_data = _status_payload()
|
|
|
|
| 537 |
|
| 538 |
active_backend = self._active_backend()
|
| 539 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 540 |
# Always expose settings UI if a settings app is available
|
| 541 |
+
# (do this AFTER loading the instance .env so status endpoint sees the right value)
|
| 542 |
self._init_settings_ui_if_needed()
|
| 543 |
|
| 544 |
# If key is still missing -> wait until provided via the settings UI
|
| 545 |
if not self._has_required_key(active_backend):
|
| 546 |
+
requirement_name = self._requirement_name(active_backend)
|
| 547 |
+
if active_backend == HF_BACKEND:
|
| 548 |
+
logger.error(
|
| 549 |
+
"%s not found. Set it in the app .env before starting the Hugging Face backend.", requirement_name
|
| 550 |
+
)
|
| 551 |
+
return
|
| 552 |
+
else:
|
| 553 |
+
logger.warning("%s not found. Open the app settings page to enter it.", requirement_name)
|
| 554 |
# Poll until the key becomes available (set via the settings UI)
|
| 555 |
try:
|
| 556 |
while not self._has_required_key(active_backend):
|
|
|
|
| 563 |
self._robot.media.start_recording()
|
| 564 |
self._robot.media.start_playing()
|
| 565 |
time.sleep(1) # give some time to the pipelines to start
|
| 566 |
+
apply_audio_startup_config(self._robot, logger=logger)
|
| 567 |
|
| 568 |
async def runner() -> None:
|
| 569 |
# Capture loop for cross-thread personality actions
|
|
|
|
| 632 |
backend = getattr(self._robot.media, "backend", None)
|
| 633 |
audio = getattr(self._robot.media, "audio", None)
|
| 634 |
if audio is not None:
|
| 635 |
+
if (
|
| 636 |
+
LOCAL_PLAYER_BACKEND is not None
|
| 637 |
+
and backend == LOCAL_PLAYER_BACKEND
|
| 638 |
+
and hasattr(audio, "clear_player")
|
| 639 |
+
and callable(audio.clear_player)
|
| 640 |
+
):
|
| 641 |
audio.clear_player()
|
| 642 |
elif (
|
| 643 |
backend == MediaBackend.WEBRTC
|
src/reachy_mini_conversation_app/conversation_handler.py
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
import asyncio
|
| 3 |
+
from abc import ABC, abstractmethod
|
| 4 |
+
from typing import TypeAlias
|
| 5 |
+
from collections.abc import Callable
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
from fastrtc import AdditionalOutputs, AsyncStreamHandler
|
| 9 |
+
from numpy.typing import NDArray
|
| 10 |
+
|
| 11 |
+
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
AudioFrame: TypeAlias = tuple[int, NDArray[np.int16]]
|
| 15 |
+
HandlerOutput: TypeAlias = AudioFrame | AdditionalOutputs | None
|
| 16 |
+
QueueItem: TypeAlias = AudioFrame | AdditionalOutputs
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class ConversationHandler(AsyncStreamHandler, ABC):
|
| 20 |
+
"""Shared app handler contract for realtime conversation backends."""
|
| 21 |
+
|
| 22 |
+
deps: ToolDependencies
|
| 23 |
+
output_queue: asyncio.Queue[QueueItem]
|
| 24 |
+
_clear_queue: Callable[[], None] | None
|
| 25 |
+
|
| 26 |
+
@abstractmethod
|
| 27 |
+
def copy(self) -> ConversationHandler:
|
| 28 |
+
"""Create a copy of the handler."""
|
| 29 |
+
...
|
| 30 |
+
|
| 31 |
+
@abstractmethod
|
| 32 |
+
async def start_up(self) -> None:
|
| 33 |
+
"""Start the realtime handler."""
|
| 34 |
+
...
|
| 35 |
+
|
| 36 |
+
@abstractmethod
|
| 37 |
+
async def shutdown(self) -> None:
|
| 38 |
+
"""Shut down the realtime handler."""
|
| 39 |
+
...
|
| 40 |
+
|
| 41 |
+
@abstractmethod
|
| 42 |
+
async def receive(self, frame: AudioFrame) -> None:
|
| 43 |
+
"""Receive an input audio frame."""
|
| 44 |
+
...
|
| 45 |
+
|
| 46 |
+
@abstractmethod
|
| 47 |
+
async def emit(self) -> HandlerOutput:
|
| 48 |
+
"""Emit the next output item."""
|
| 49 |
+
...
|
| 50 |
+
|
| 51 |
+
@abstractmethod
|
| 52 |
+
async def apply_personality(self, profile: str | None) -> str:
|
| 53 |
+
"""Apply a personality profile."""
|
| 54 |
+
...
|
| 55 |
+
|
| 56 |
+
@abstractmethod
|
| 57 |
+
async def get_available_voices(self) -> list[str]:
|
| 58 |
+
"""Return voices available for the active backend."""
|
| 59 |
+
...
|
| 60 |
+
|
| 61 |
+
@abstractmethod
|
| 62 |
+
def get_current_voice(self) -> str:
|
| 63 |
+
"""Return the current voice."""
|
| 64 |
+
...
|
| 65 |
+
|
| 66 |
+
@abstractmethod
|
| 67 |
+
async def change_voice(self, voice: str) -> str:
|
| 68 |
+
"""Change the current voice."""
|
| 69 |
+
...
|
src/reachy_mini_conversation_app/gemini_live.py
CHANGED
|
@@ -20,7 +20,7 @@ from datetime import datetime
|
|
| 20 |
import numpy as np
|
| 21 |
import gradio as gr
|
| 22 |
from google import genai
|
| 23 |
-
from fastrtc import AdditionalOutputs,
|
| 24 |
from google.genai import types
|
| 25 |
from numpy.typing import NDArray
|
| 26 |
from scipy.signal import resample
|
|
@@ -34,8 +34,9 @@ from reachy_mini_conversation_app.config import (
|
|
| 34 |
from reachy_mini_conversation_app.prompts import get_session_voice, get_session_instructions
|
| 35 |
from reachy_mini_conversation_app.tools.core_tools import (
|
| 36 |
ToolDependencies,
|
| 37 |
-
|
| 38 |
)
|
|
|
|
| 39 |
from reachy_mini_conversation_app.camera_frame_encoding import encode_bgr_frame_as_jpeg
|
| 40 |
from reachy_mini_conversation_app.tools.background_tool_manager import (
|
| 41 |
ToolCallRoutine,
|
|
@@ -121,10 +122,32 @@ def _resolve_gemini_voice(profile_voice: str) -> str:
|
|
| 121 |
return voice_map.get(profile_voice.lower(), DEFAULT_VOICE_BY_BACKEND[GEMINI_BACKEND])
|
| 122 |
|
| 123 |
|
| 124 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 125 |
"""Gemini Live API handler for fastrtc Stream."""
|
| 126 |
|
| 127 |
-
def __init__(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 128 |
"""Initialize the handler."""
|
| 129 |
super().__init__(
|
| 130 |
expected_layout="mono",
|
|
@@ -135,7 +158,7 @@ class GeminiLiveHandler(AsyncStreamHandler):
|
|
| 135 |
self.deps = deps
|
| 136 |
self.gradio_mode = gradio_mode
|
| 137 |
self.instance_path = instance_path
|
| 138 |
-
self._voice_override: str | None =
|
| 139 |
|
| 140 |
self.session: Any = None # google.genai live session
|
| 141 |
self.output_queue: "asyncio.Queue[Tuple[int, NDArray[np.int16]] | AdditionalOutputs]" = asyncio.Queue()
|
|
@@ -162,7 +185,12 @@ class GeminiLiveHandler(AsyncStreamHandler):
|
|
| 162 |
|
| 163 |
def copy(self) -> "GeminiLiveHandler":
|
| 164 |
"""Create a copy of the handler."""
|
| 165 |
-
return GeminiLiveHandler(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 166 |
|
| 167 |
def _set_listening_state(self, listening: bool) -> None:
|
| 168 |
"""Avoid queueing redundant listening-state updates."""
|
|
@@ -217,7 +245,6 @@ class GeminiLiveHandler(AsyncStreamHandler):
|
|
| 217 |
from reachy_mini_conversation_app.config import set_custom_profile
|
| 218 |
|
| 219 |
set_custom_profile(profile)
|
| 220 |
-
self._voice_override = None
|
| 221 |
logger.info("Set custom profile to %r", profile)
|
| 222 |
|
| 223 |
try:
|
|
@@ -341,7 +368,11 @@ class GeminiLiveHandler(AsyncStreamHandler):
|
|
| 341 |
voice = _resolve_gemini_voice(self._voice_override or get_session_voice())
|
| 342 |
|
| 343 |
# Convert OpenAI-style tool specs to Gemini function declarations
|
| 344 |
-
tool_specs =
|
|
|
|
|
|
|
|
|
|
|
|
|
| 345 |
function_declarations = _openai_tool_specs_to_gemini(tool_specs)
|
| 346 |
|
| 347 |
tools_config: List[Dict[str, Any]] = []
|
|
@@ -700,7 +731,7 @@ class GeminiLiveHandler(AsyncStreamHandler):
|
|
| 700 |
timestamp_msg = (
|
| 701 |
f"[Idle time update: {self.format_timestamp()} - No activity for {idle_duration:.1f}s] "
|
| 702 |
"You've been idle for a while. Feel free to get creative - dance, show an emotion, "
|
| 703 |
-
"look around,
|
| 704 |
)
|
| 705 |
if not self.session:
|
| 706 |
logger.debug("No session, cannot send idle signal")
|
|
|
|
| 20 |
import numpy as np
|
| 21 |
import gradio as gr
|
| 22 |
from google import genai
|
| 23 |
+
from fastrtc import AdditionalOutputs, wait_for_item, audio_to_int16
|
| 24 |
from google.genai import types
|
| 25 |
from numpy.typing import NDArray
|
| 26 |
from scipy.signal import resample
|
|
|
|
| 34 |
from reachy_mini_conversation_app.prompts import get_session_voice, get_session_instructions
|
| 35 |
from reachy_mini_conversation_app.tools.core_tools import (
|
| 36 |
ToolDependencies,
|
| 37 |
+
get_active_tool_specs,
|
| 38 |
)
|
| 39 |
+
from reachy_mini_conversation_app.conversation_handler import ConversationHandler
|
| 40 |
from reachy_mini_conversation_app.camera_frame_encoding import encode_bgr_frame_as_jpeg
|
| 41 |
from reachy_mini_conversation_app.tools.background_tool_manager import (
|
| 42 |
ToolCallRoutine,
|
|
|
|
| 122 |
return voice_map.get(profile_voice.lower(), DEFAULT_VOICE_BY_BACKEND[GEMINI_BACKEND])
|
| 123 |
|
| 124 |
|
| 125 |
+
def _resolve_gemini_startup_voice(voice: str | None) -> str | None:
|
| 126 |
+
"""Return a valid persisted Gemini startup voice or None."""
|
| 127 |
+
if voice is None:
|
| 128 |
+
return None
|
| 129 |
+
|
| 130 |
+
voice_map = {candidate.lower(): candidate for candidate in GEMINI_AVAILABLE_VOICES}
|
| 131 |
+
resolved = voice_map.get(voice.lower())
|
| 132 |
+
if resolved is None:
|
| 133 |
+
logger.warning(
|
| 134 |
+
"Ignoring persisted Gemini startup voice %r; expected one of %s",
|
| 135 |
+
voice,
|
| 136 |
+
GEMINI_AVAILABLE_VOICES,
|
| 137 |
+
)
|
| 138 |
+
return resolved
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
class GeminiLiveHandler(ConversationHandler):
|
| 142 |
"""Gemini Live API handler for fastrtc Stream."""
|
| 143 |
|
| 144 |
+
def __init__(
|
| 145 |
+
self,
|
| 146 |
+
deps: ToolDependencies,
|
| 147 |
+
gradio_mode: bool = False,
|
| 148 |
+
instance_path: Optional[str] = None,
|
| 149 |
+
startup_voice: Optional[str] = None,
|
| 150 |
+
):
|
| 151 |
"""Initialize the handler."""
|
| 152 |
super().__init__(
|
| 153 |
expected_layout="mono",
|
|
|
|
| 158 |
self.deps = deps
|
| 159 |
self.gradio_mode = gradio_mode
|
| 160 |
self.instance_path = instance_path
|
| 161 |
+
self._voice_override: str | None = _resolve_gemini_startup_voice(startup_voice)
|
| 162 |
|
| 163 |
self.session: Any = None # google.genai live session
|
| 164 |
self.output_queue: "asyncio.Queue[Tuple[int, NDArray[np.int16]] | AdditionalOutputs]" = asyncio.Queue()
|
|
|
|
| 185 |
|
| 186 |
def copy(self) -> "GeminiLiveHandler":
|
| 187 |
"""Create a copy of the handler."""
|
| 188 |
+
return GeminiLiveHandler(
|
| 189 |
+
self.deps,
|
| 190 |
+
self.gradio_mode,
|
| 191 |
+
self.instance_path,
|
| 192 |
+
startup_voice=self._voice_override,
|
| 193 |
+
)
|
| 194 |
|
| 195 |
def _set_listening_state(self, listening: bool) -> None:
|
| 196 |
"""Avoid queueing redundant listening-state updates."""
|
|
|
|
| 245 |
from reachy_mini_conversation_app.config import set_custom_profile
|
| 246 |
|
| 247 |
set_custom_profile(profile)
|
|
|
|
| 248 |
logger.info("Set custom profile to %r", profile)
|
| 249 |
|
| 250 |
try:
|
|
|
|
| 368 |
voice = _resolve_gemini_voice(self._voice_override or get_session_voice())
|
| 369 |
|
| 370 |
# Convert OpenAI-style tool specs to Gemini function declarations
|
| 371 |
+
tool_specs = get_active_tool_specs(self.deps)
|
| 372 |
+
logger.info(
|
| 373 |
+
"Tools to be used in conversation: %s",
|
| 374 |
+
[tool["name"] for tool in tool_specs],
|
| 375 |
+
)
|
| 376 |
function_declarations = _openai_tool_specs_to_gemini(tool_specs)
|
| 377 |
|
| 378 |
tools_config: List[Dict[str, Any]] = []
|
|
|
|
| 731 |
timestamp_msg = (
|
| 732 |
f"[Idle time update: {self.format_timestamp()} - No activity for {idle_duration:.1f}s] "
|
| 733 |
"You've been idle for a while. Feel free to get creative - dance, show an emotion, "
|
| 734 |
+
"look around, call idle_do_nothing to stay still and silent, or just be yourself!"
|
| 735 |
)
|
| 736 |
if not self.session:
|
| 737 |
logger.debug("No session, cannot send idle signal")
|
src/reachy_mini_conversation_app/gradio_personality.py
CHANGED
|
@@ -79,6 +79,44 @@ class PersonalityUI:
|
|
| 79 |
except Exception as e:
|
| 80 |
return f"Could not load instructions: {e}"
|
| 81 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 82 |
@staticmethod
|
| 83 |
def _sanitize_name(name: str) -> str:
|
| 84 |
import re
|
|
@@ -101,6 +139,10 @@ class PersonalityUI:
|
|
| 101 |
current_value = config.REACHY_MINI_CUSTOM_PROFILE or self.DEFAULT_OPTION
|
| 102 |
dropdown_label = "Select personality"
|
| 103 |
dropdown_choices = [self.DEFAULT_OPTION, *(self._list_personalities())]
|
|
|
|
|
|
|
|
|
|
|
|
|
| 104 |
|
| 105 |
self.personalities_dropdown = gr.Dropdown(
|
| 106 |
label=dropdown_label,
|
|
@@ -113,7 +155,12 @@ class PersonalityUI:
|
|
| 113 |
self.preview_md = gr.Markdown(value=self._read_instructions_for(current_value))
|
| 114 |
self.person_name_tb = gr.Textbox(label="Personality name", interactive=not is_locked)
|
| 115 |
self.person_instr_ta = gr.TextArea(label="Personality instructions", lines=10, interactive=not is_locked)
|
| 116 |
-
self.tools_txt_ta = gr.TextArea(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 117 |
self.voice_dropdown = gr.Dropdown(
|
| 118 |
label="Voice",
|
| 119 |
choices=get_available_voices_for_backend(),
|
|
@@ -122,7 +169,10 @@ class PersonalityUI:
|
|
| 122 |
)
|
| 123 |
self.new_personality_btn = gr.Button("New personality", interactive=not is_locked)
|
| 124 |
self.available_tools_cg = gr.CheckboxGroup(
|
| 125 |
-
label="Available tools (helper)",
|
|
|
|
|
|
|
|
|
|
| 126 |
)
|
| 127 |
self.save_btn = gr.Button("Save personality (instructions + tools)", interactive=not is_locked)
|
| 128 |
|
|
@@ -183,43 +233,12 @@ class PersonalityUI:
|
|
| 183 |
value=get_default_voice_for_backend(),
|
| 184 |
)
|
| 185 |
|
| 186 |
-
def _available_tools_for(selected: str) -> tuple[list[str], list[str]]:
|
| 187 |
-
shared: list[str] = []
|
| 188 |
-
try:
|
| 189 |
-
for py in self._tools_dir.glob("*.py"):
|
| 190 |
-
if py.stem in {"__init__", "core_tools"}:
|
| 191 |
-
continue
|
| 192 |
-
shared.append(py.stem)
|
| 193 |
-
except Exception:
|
| 194 |
-
pass
|
| 195 |
-
local: list[str] = []
|
| 196 |
-
try:
|
| 197 |
-
if selected != self.DEFAULT_OPTION:
|
| 198 |
-
for py in (self._profiles_root / selected).glob("*.py"):
|
| 199 |
-
local.append(py.stem)
|
| 200 |
-
except Exception:
|
| 201 |
-
pass
|
| 202 |
-
return sorted(shared), sorted(local)
|
| 203 |
-
|
| 204 |
-
def _parse_enabled_tools(text: str) -> list[str]:
|
| 205 |
-
enabled: list[str] = []
|
| 206 |
-
for line in text.splitlines():
|
| 207 |
-
s = line.strip()
|
| 208 |
-
if not s or s.startswith("#"):
|
| 209 |
-
continue
|
| 210 |
-
enabled.append(s)
|
| 211 |
-
return enabled
|
| 212 |
-
|
| 213 |
def _load_profile_for_edit(selected: str) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any], str]:
|
| 214 |
instr = self._read_instructions_for(selected)
|
| 215 |
-
tools_txt =
|
| 216 |
-
|
| 217 |
-
tp = self._resolve_profile_dir(selected) / "tools.txt"
|
| 218 |
-
if tp.exists():
|
| 219 |
-
tools_txt = tp.read_text(encoding="utf-8")
|
| 220 |
-
shared, local = _available_tools_for(selected)
|
| 221 |
all_tools = sorted(set(shared + local))
|
| 222 |
-
enabled = _parse_enabled_tools(tools_txt)
|
| 223 |
status_text = f"Loaded profile '{selected}'."
|
| 224 |
return (
|
| 225 |
gr.update(value=instr),
|
|
@@ -239,7 +258,7 @@ class PersonalityUI:
|
|
| 239 |
gr.update(value=""),
|
| 240 |
gr.update(value=instr_val),
|
| 241 |
gr.update(value=tools_txt_val),
|
| 242 |
-
gr.update(choices=sorted(_available_tools_for(self.DEFAULT_OPTION)[0]), value=[]),
|
| 243 |
"Fill in a name, instructions and (optional) tools, then Save.",
|
| 244 |
gr.update(value=get_default_voice_for_backend()),
|
| 245 |
)
|
|
|
|
| 79 |
except Exception as e:
|
| 80 |
return f"Could not load instructions: {e}"
|
| 81 |
|
| 82 |
+
def _read_tools_for(self, name: str) -> str:
|
| 83 |
+
try:
|
| 84 |
+
profile_name = "default" if name == self.DEFAULT_OPTION else name
|
| 85 |
+
target = self._resolve_profile_dir(profile_name) / "tools.txt"
|
| 86 |
+
if target.exists():
|
| 87 |
+
return target.read_text(encoding="utf-8")
|
| 88 |
+
except Exception:
|
| 89 |
+
pass
|
| 90 |
+
return ""
|
| 91 |
+
|
| 92 |
+
def _available_tools_for(self, selected: str) -> tuple[list[str], list[str]]:
|
| 93 |
+
shared: list[str] = []
|
| 94 |
+
try:
|
| 95 |
+
for py in self._tools_dir.glob("*.py"):
|
| 96 |
+
if py.stem in {"__init__", "core_tools"}:
|
| 97 |
+
continue
|
| 98 |
+
shared.append(py.stem)
|
| 99 |
+
except Exception:
|
| 100 |
+
pass
|
| 101 |
+
local: list[str] = []
|
| 102 |
+
try:
|
| 103 |
+
if selected != self.DEFAULT_OPTION:
|
| 104 |
+
for py in (self._profiles_root / selected).glob("*.py"):
|
| 105 |
+
local.append(py.stem)
|
| 106 |
+
except Exception:
|
| 107 |
+
pass
|
| 108 |
+
return sorted(shared), sorted(local)
|
| 109 |
+
|
| 110 |
+
@staticmethod
|
| 111 |
+
def _parse_enabled_tools(text: str) -> list[str]:
|
| 112 |
+
enabled: list[str] = []
|
| 113 |
+
for line in text.splitlines():
|
| 114 |
+
s = line.strip()
|
| 115 |
+
if not s or s.startswith("#"):
|
| 116 |
+
continue
|
| 117 |
+
enabled.append(s)
|
| 118 |
+
return enabled
|
| 119 |
+
|
| 120 |
@staticmethod
|
| 121 |
def _sanitize_name(name: str) -> str:
|
| 122 |
import re
|
|
|
|
| 139 |
current_value = config.REACHY_MINI_CUSTOM_PROFILE or self.DEFAULT_OPTION
|
| 140 |
dropdown_label = "Select personality"
|
| 141 |
dropdown_choices = [self.DEFAULT_OPTION, *(self._list_personalities())]
|
| 142 |
+
initial_tools_txt = self._read_tools_for(current_value)
|
| 143 |
+
shared_tools, local_tools = self._available_tools_for(current_value)
|
| 144 |
+
initial_available_tools = sorted(set(shared_tools + local_tools))
|
| 145 |
+
initial_enabled_tools = self._parse_enabled_tools(initial_tools_txt)
|
| 146 |
|
| 147 |
self.personalities_dropdown = gr.Dropdown(
|
| 148 |
label=dropdown_label,
|
|
|
|
| 155 |
self.preview_md = gr.Markdown(value=self._read_instructions_for(current_value))
|
| 156 |
self.person_name_tb = gr.Textbox(label="Personality name", interactive=not is_locked)
|
| 157 |
self.person_instr_ta = gr.TextArea(label="Personality instructions", lines=10, interactive=not is_locked)
|
| 158 |
+
self.tools_txt_ta = gr.TextArea(
|
| 159 |
+
label="tools.txt",
|
| 160 |
+
value=initial_tools_txt,
|
| 161 |
+
lines=10,
|
| 162 |
+
interactive=not is_locked,
|
| 163 |
+
)
|
| 164 |
self.voice_dropdown = gr.Dropdown(
|
| 165 |
label="Voice",
|
| 166 |
choices=get_available_voices_for_backend(),
|
|
|
|
| 169 |
)
|
| 170 |
self.new_personality_btn = gr.Button("New personality", interactive=not is_locked)
|
| 171 |
self.available_tools_cg = gr.CheckboxGroup(
|
| 172 |
+
label="Available tools (helper)",
|
| 173 |
+
choices=initial_available_tools,
|
| 174 |
+
value=initial_enabled_tools,
|
| 175 |
+
interactive=not is_locked,
|
| 176 |
)
|
| 177 |
self.save_btn = gr.Button("Save personality (instructions + tools)", interactive=not is_locked)
|
| 178 |
|
|
|
|
| 233 |
value=get_default_voice_for_backend(),
|
| 234 |
)
|
| 235 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 236 |
def _load_profile_for_edit(selected: str) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any], str]:
|
| 237 |
instr = self._read_instructions_for(selected)
|
| 238 |
+
tools_txt = self._read_tools_for(selected)
|
| 239 |
+
shared, local = self._available_tools_for(selected)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 240 |
all_tools = sorted(set(shared + local))
|
| 241 |
+
enabled = self._parse_enabled_tools(tools_txt)
|
| 242 |
status_text = f"Loaded profile '{selected}'."
|
| 243 |
return (
|
| 244 |
gr.update(value=instr),
|
|
|
|
| 258 |
gr.update(value=""),
|
| 259 |
gr.update(value=instr_val),
|
| 260 |
gr.update(value=tools_txt_val),
|
| 261 |
+
gr.update(choices=sorted(self._available_tools_for(self.DEFAULT_OPTION)[0]), value=[]),
|
| 262 |
"Fill in a name, instructions and (optional) tools, then Save.",
|
| 263 |
gr.update(value=get_default_voice_for_backend()),
|
| 264 |
)
|
src/reachy_mini_conversation_app/headless_personality.py
CHANGED
|
@@ -11,7 +11,7 @@ from __future__ import annotations
|
|
| 11 |
from typing import List
|
| 12 |
from pathlib import Path
|
| 13 |
|
| 14 |
-
from .config import DEFAULT_PROFILES_DIRECTORY
|
| 15 |
|
| 16 |
|
| 17 |
DEFAULT_OPTION = "(built-in default)"
|
|
@@ -76,6 +76,16 @@ def read_instructions_for(name: str) -> str:
|
|
| 76 |
return f"Could not load instructions: {e}"
|
| 77 |
|
| 78 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 79 |
def available_tools_for(selected: str) -> List[str]:
|
| 80 |
"""List available tool modules for the given profile selection."""
|
| 81 |
shared: List[str] = []
|
|
@@ -96,9 +106,10 @@ def available_tools_for(selected: str) -> List[str]:
|
|
| 96 |
return sorted(set(shared + local))
|
| 97 |
|
| 98 |
|
| 99 |
-
def _write_profile(name_s: str, instructions: str, tools_text: str, voice: str =
|
|
|
|
| 100 |
target_dir = _profiles_root() / "user_personalities" / name_s
|
| 101 |
target_dir.mkdir(parents=True, exist_ok=True)
|
| 102 |
(target_dir / "instructions.txt").write_text(instructions.strip() + "\n", encoding="utf-8")
|
| 103 |
(target_dir / "tools.txt").write_text((tools_text or "").strip() + "\n", encoding="utf-8")
|
| 104 |
-
(target_dir / "voice.txt").write_text((voice or
|
|
|
|
| 11 |
from typing import List
|
| 12 |
from pathlib import Path
|
| 13 |
|
| 14 |
+
from .config import DEFAULT_PROFILES_DIRECTORY, get_default_voice_for_backend
|
| 15 |
|
| 16 |
|
| 17 |
DEFAULT_OPTION = "(built-in default)"
|
|
|
|
| 76 |
return f"Could not load instructions: {e}"
|
| 77 |
|
| 78 |
|
| 79 |
+
def read_tools_for(name: str) -> str:
|
| 80 |
+
"""Read the tools.txt content for the given profile name."""
|
| 81 |
+
try:
|
| 82 |
+
profile_name = "default" if name == DEFAULT_OPTION else name
|
| 83 |
+
target = resolve_profile_dir(profile_name) / "tools.txt"
|
| 84 |
+
return target.read_text(encoding="utf-8") if target.exists() else ""
|
| 85 |
+
except Exception:
|
| 86 |
+
return ""
|
| 87 |
+
|
| 88 |
+
|
| 89 |
def available_tools_for(selected: str) -> List[str]:
|
| 90 |
"""List available tool modules for the given profile selection."""
|
| 91 |
shared: List[str] = []
|
|
|
|
| 106 |
return sorted(set(shared + local))
|
| 107 |
|
| 108 |
|
| 109 |
+
def _write_profile(name_s: str, instructions: str, tools_text: str, voice: str | None = None) -> None:
|
| 110 |
+
default_voice = get_default_voice_for_backend()
|
| 111 |
target_dir = _profiles_root() / "user_personalities" / name_s
|
| 112 |
target_dir.mkdir(parents=True, exist_ok=True)
|
| 113 |
(target_dir / "instructions.txt").write_text(instructions.strip() + "\n", encoding="utf-8")
|
| 114 |
(target_dir / "tools.txt").write_text((tools_text or "").strip() + "\n", encoding="utf-8")
|
| 115 |
+
(target_dir / "voice.txt").write_text((voice or default_voice).strip() + "\n", encoding="utf-8")
|
src/reachy_mini_conversation_app/headless_personality_ui.py
CHANGED
|
@@ -9,7 +9,7 @@ callable to avoid cross-thread issues.
|
|
| 9 |
from __future__ import annotations
|
| 10 |
import asyncio
|
| 11 |
import logging
|
| 12 |
-
from typing import
|
| 13 |
|
| 14 |
from fastapi import Query, FastAPI, Request
|
| 15 |
|
|
@@ -19,15 +19,12 @@ from .config import (
|
|
| 19 |
get_default_voice_for_backend,
|
| 20 |
get_available_voices_for_backend,
|
| 21 |
)
|
| 22 |
-
from .
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
if TYPE_CHECKING:
|
| 26 |
-
from .gemini_live import GeminiLiveHandler
|
| 27 |
from .headless_personality import (
|
| 28 |
DEFAULT_OPTION,
|
| 29 |
_sanitize_name,
|
| 30 |
_write_profile,
|
|
|
|
| 31 |
list_personalities,
|
| 32 |
available_tools_for,
|
| 33 |
resolve_profile_dir,
|
|
@@ -40,10 +37,10 @@ logger = logging.getLogger(__name__)
|
|
| 40 |
|
| 41 |
def mount_personality_routes(
|
| 42 |
app: FastAPI,
|
| 43 |
-
handler:
|
| 44 |
get_loop: Callable[[], asyncio.AbstractEventLoop | None],
|
| 45 |
*,
|
| 46 |
-
persist_personality: Callable[[Optional[str]], None] | None = None,
|
| 47 |
get_persisted_personality: Callable[[], Optional[str]] | None = None,
|
| 48 |
) -> None:
|
| 49 |
"""Register personality management endpoints on a FastAPI app."""
|
|
@@ -92,14 +89,11 @@ def mount_personality_routes(
|
|
| 92 |
@app.get("/personalities/load")
|
| 93 |
def _load(name: str) -> dict: # type: ignore
|
| 94 |
instr = read_instructions_for(name)
|
| 95 |
-
tools_txt =
|
| 96 |
voice = get_default_voice_for_backend()
|
| 97 |
uses_default_voice = True
|
| 98 |
if name != DEFAULT_OPTION:
|
| 99 |
pdir = resolve_profile_dir(name)
|
| 100 |
-
tp = pdir / "tools.txt"
|
| 101 |
-
if tp.exists():
|
| 102 |
-
tools_txt = tp.read_text(encoding="utf-8")
|
| 103 |
vf = pdir / "voice.txt"
|
| 104 |
if vf.exists():
|
| 105 |
v = vf.read_text(encoding="utf-8").strip()
|
|
@@ -258,19 +252,21 @@ def mount_personality_routes(
|
|
| 258 |
if not sel_name:
|
| 259 |
sel_name = DEFAULT_OPTION
|
| 260 |
|
| 261 |
-
async def _do_apply() -> str:
|
| 262 |
sel = None if sel_name == DEFAULT_OPTION else sel_name
|
| 263 |
status = await handler.apply_personality(sel)
|
| 264 |
-
|
|
|
|
|
|
|
| 265 |
|
| 266 |
try:
|
| 267 |
logger.info("Headless apply: requested name=%r", sel_name)
|
| 268 |
fut = asyncio.run_coroutine_threadsafe(_do_apply(), loop)
|
| 269 |
-
status = fut.result(timeout=10)
|
| 270 |
persisted_choice = _startup_choice()
|
| 271 |
if persist_flag and persist_personality is not None:
|
| 272 |
try:
|
| 273 |
-
persist_personality(None if sel_name == DEFAULT_OPTION else sel_name)
|
| 274 |
persisted_choice = _startup_choice()
|
| 275 |
except Exception as e:
|
| 276 |
logger.warning("Failed to persist startup personality: %s", e)
|
|
|
|
| 9 |
from __future__ import annotations
|
| 10 |
import asyncio
|
| 11 |
import logging
|
| 12 |
+
from typing import Any, Callable, Optional
|
| 13 |
|
| 14 |
from fastapi import Query, FastAPI, Request
|
| 15 |
|
|
|
|
| 19 |
get_default_voice_for_backend,
|
| 20 |
get_available_voices_for_backend,
|
| 21 |
)
|
| 22 |
+
from .conversation_handler import ConversationHandler
|
|
|
|
|
|
|
|
|
|
|
|
|
| 23 |
from .headless_personality import (
|
| 24 |
DEFAULT_OPTION,
|
| 25 |
_sanitize_name,
|
| 26 |
_write_profile,
|
| 27 |
+
read_tools_for,
|
| 28 |
list_personalities,
|
| 29 |
available_tools_for,
|
| 30 |
resolve_profile_dir,
|
|
|
|
| 37 |
|
| 38 |
def mount_personality_routes(
|
| 39 |
app: FastAPI,
|
| 40 |
+
handler: ConversationHandler,
|
| 41 |
get_loop: Callable[[], asyncio.AbstractEventLoop | None],
|
| 42 |
*,
|
| 43 |
+
persist_personality: Callable[[Optional[str], Optional[str]], None] | None = None,
|
| 44 |
get_persisted_personality: Callable[[], Optional[str]] | None = None,
|
| 45 |
) -> None:
|
| 46 |
"""Register personality management endpoints on a FastAPI app."""
|
|
|
|
| 89 |
@app.get("/personalities/load")
|
| 90 |
def _load(name: str) -> dict: # type: ignore
|
| 91 |
instr = read_instructions_for(name)
|
| 92 |
+
tools_txt = read_tools_for(name)
|
| 93 |
voice = get_default_voice_for_backend()
|
| 94 |
uses_default_voice = True
|
| 95 |
if name != DEFAULT_OPTION:
|
| 96 |
pdir = resolve_profile_dir(name)
|
|
|
|
|
|
|
|
|
|
| 97 |
vf = pdir / "voice.txt"
|
| 98 |
if vf.exists():
|
| 99 |
v = vf.read_text(encoding="utf-8").strip()
|
|
|
|
| 252 |
if not sel_name:
|
| 253 |
sel_name = DEFAULT_OPTION
|
| 254 |
|
| 255 |
+
async def _do_apply() -> tuple[str, Optional[str]]:
|
| 256 |
sel = None if sel_name == DEFAULT_OPTION else sel_name
|
| 257 |
status = await handler.apply_personality(sel)
|
| 258 |
+
get_current_voice = getattr(handler, "get_current_voice", None)
|
| 259 |
+
voice_override = get_current_voice() if callable(get_current_voice) else None
|
| 260 |
+
return status, voice_override
|
| 261 |
|
| 262 |
try:
|
| 263 |
logger.info("Headless apply: requested name=%r", sel_name)
|
| 264 |
fut = asyncio.run_coroutine_threadsafe(_do_apply(), loop)
|
| 265 |
+
status, voice_override = fut.result(timeout=10)
|
| 266 |
persisted_choice = _startup_choice()
|
| 267 |
if persist_flag and persist_personality is not None:
|
| 268 |
try:
|
| 269 |
+
persist_personality(None if sel_name == DEFAULT_OPTION else sel_name, voice_override)
|
| 270 |
persisted_choice = _startup_choice()
|
| 271 |
except Exception as e:
|
| 272 |
logger.warning("Failed to persist startup personality: %s", e)
|
src/reachy_mini_conversation_app/huggingface_realtime.py
ADDED
|
@@ -0,0 +1,160 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from typing import Any
|
| 3 |
+
|
| 4 |
+
import httpx
|
| 5 |
+
from openai import AsyncOpenAI
|
| 6 |
+
from typing_extensions import Literal, TypedDict
|
| 7 |
+
from openai.types.realtime import (
|
| 8 |
+
AudioTranscriptionParam,
|
| 9 |
+
RealtimeAudioConfigParam,
|
| 10 |
+
RealtimeAudioConfigInputParam,
|
| 11 |
+
RealtimeAudioConfigOutputParam,
|
| 12 |
+
RealtimeSessionCreateRequestParam,
|
| 13 |
+
)
|
| 14 |
+
from openai.types.realtime.realtime_audio_input_turn_detection_param import ServerVad
|
| 15 |
+
|
| 16 |
+
from reachy_mini_conversation_app.config import (
|
| 17 |
+
HF_BACKEND,
|
| 18 |
+
HF_LOCAL_CONNECTION_MODE,
|
| 19 |
+
config,
|
| 20 |
+
get_hf_direct_ws_url,
|
| 21 |
+
parse_hf_realtime_url,
|
| 22 |
+
get_hf_connection_selection,
|
| 23 |
+
)
|
| 24 |
+
from reachy_mini_conversation_app.prompts import get_session_voice, get_session_instructions
|
| 25 |
+
from reachy_mini_conversation_app.base_realtime import (
|
| 26 |
+
BaseRealtimeHandler,
|
| 27 |
+
InputTranscriptChunksByItem,
|
| 28 |
+
to_realtime_tools_config,
|
| 29 |
+
)
|
| 30 |
+
from reachy_mini_conversation_app.tools.core_tools import get_active_tool_specs
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
logger = logging.getLogger(__name__)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def _build_openai_compatible_client_from_realtime_url(
|
| 37 |
+
realtime_url: str,
|
| 38 |
+
bearer_token: str | None,
|
| 39 |
+
) -> tuple[AsyncOpenAI, dict[str, str]]:
|
| 40 |
+
"""Build an OpenAI-compatible realtime client from a direct websocket/base URL."""
|
| 41 |
+
parsed = parse_hf_realtime_url(realtime_url)
|
| 42 |
+
client = AsyncOpenAI(
|
| 43 |
+
api_key=bearer_token or "DUMMY",
|
| 44 |
+
base_url=parsed.base_url,
|
| 45 |
+
websocket_base_url=parsed.websocket_base_url,
|
| 46 |
+
)
|
| 47 |
+
return client, parsed.connect_query
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
class HFNativeRateAudioPCM(TypedDict):
|
| 51 |
+
"""Hugging Face extension for native-rate PCM audio."""
|
| 52 |
+
|
| 53 |
+
type: Literal["audio/pcm"]
|
| 54 |
+
rate: None
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def _native_rate_audio_pcm() -> HFNativeRateAudioPCM:
|
| 58 |
+
"""Return the Hugging Face native-rate PCM config."""
|
| 59 |
+
return {"type": "audio/pcm", "rate": None}
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
class HuggingFaceRealtimeHandler(BaseRealtimeHandler):
|
| 63 |
+
"""Realtime handler for Hugging Face endpoints."""
|
| 64 |
+
|
| 65 |
+
BACKEND_PROVIDER = HF_BACKEND
|
| 66 |
+
SAMPLE_RATE = 16000
|
| 67 |
+
REFRESH_CLIENT_ON_RECONNECT = True
|
| 68 |
+
AUDIO_INPUT_COST_PER_1M = 0.0
|
| 69 |
+
AUDIO_OUTPUT_COST_PER_1M = 0.0
|
| 70 |
+
TEXT_INPUT_COST_PER_1M = 0.0
|
| 71 |
+
TEXT_OUTPUT_COST_PER_1M = 0.0
|
| 72 |
+
IMAGE_INPUT_COST_PER_1M = 0.0
|
| 73 |
+
|
| 74 |
+
def _get_session_instructions(self) -> str:
|
| 75 |
+
"""Return Hugging Face session instructions."""
|
| 76 |
+
return get_session_instructions()
|
| 77 |
+
|
| 78 |
+
def _get_session_voice(self, default: str | None = None) -> str:
|
| 79 |
+
"""Return the configured Hugging Face session voice."""
|
| 80 |
+
return get_session_voice(default)
|
| 81 |
+
|
| 82 |
+
def _get_active_tool_specs(self) -> list[dict[str, Any]]:
|
| 83 |
+
"""Return active tool specs for the current session dependencies."""
|
| 84 |
+
return get_active_tool_specs(self.deps)
|
| 85 |
+
|
| 86 |
+
def _get_session_config(self, tool_specs: list[dict[str, Any]]) -> RealtimeSessionCreateRequestParam:
|
| 87 |
+
"""Return the Hugging Face OpenAI-compatible session config."""
|
| 88 |
+
return RealtimeSessionCreateRequestParam(
|
| 89 |
+
type="realtime",
|
| 90 |
+
instructions=self._get_session_instructions(),
|
| 91 |
+
audio=RealtimeAudioConfigParam(
|
| 92 |
+
input=RealtimeAudioConfigInputParam(
|
| 93 |
+
# The OpenAI SDK type only includes 24 kHz PCM, but the HF
|
| 94 |
+
# compatible server uses rate=None for native 16 kHz mode.
|
| 95 |
+
format=_native_rate_audio_pcm(), # type: ignore[typeddict-item]
|
| 96 |
+
transcription=AudioTranscriptionParam(model="gpt-4o-transcribe", language="en"),
|
| 97 |
+
turn_detection=ServerVad(type="server_vad", interrupt_response=True),
|
| 98 |
+
),
|
| 99 |
+
output=RealtimeAudioConfigOutputParam(
|
| 100 |
+
format=_native_rate_audio_pcm(), # type: ignore[typeddict-item]
|
| 101 |
+
voice=self.get_current_voice(),
|
| 102 |
+
),
|
| 103 |
+
),
|
| 104 |
+
tools=to_realtime_tools_config(tool_specs),
|
| 105 |
+
tool_choice="auto",
|
| 106 |
+
)
|
| 107 |
+
|
| 108 |
+
def _record_partial_transcript_delta(
|
| 109 |
+
self,
|
| 110 |
+
input_transcript: InputTranscriptChunksByItem,
|
| 111 |
+
item_id: str,
|
| 112 |
+
delta: str,
|
| 113 |
+
) -> None:
|
| 114 |
+
"""Record a Hugging Face partial transcript snapshot."""
|
| 115 |
+
input_transcript.item_id = item_id
|
| 116 |
+
input_transcript.deltas = [delta]
|
| 117 |
+
|
| 118 |
+
async def _build_realtime_client(self) -> AsyncOpenAI:
|
| 119 |
+
"""Build the Hugging Face OpenAI-compatible realtime client."""
|
| 120 |
+
bearer_token = (config.HF_TOKEN or "").strip()
|
| 121 |
+
connection_selection = get_hf_connection_selection()
|
| 122 |
+
direct_realtime_url = get_hf_direct_ws_url()
|
| 123 |
+
if connection_selection.mode == HF_LOCAL_CONNECTION_MODE:
|
| 124 |
+
if not direct_realtime_url:
|
| 125 |
+
raise RuntimeError("HF_REALTIME_WS_URL must be set when HF_REALTIME_CONNECTION_MODE=local")
|
| 126 |
+
client, connect_query = _build_openai_compatible_client_from_realtime_url(
|
| 127 |
+
direct_realtime_url,
|
| 128 |
+
bearer_token,
|
| 129 |
+
)
|
| 130 |
+
self._realtime_connect_query = connect_query
|
| 131 |
+
logger.info("Using direct Hugging Face realtime endpoint %s", direct_realtime_url)
|
| 132 |
+
return client
|
| 133 |
+
|
| 134 |
+
session_url = connection_selection.session_url
|
| 135 |
+
if not session_url:
|
| 136 |
+
raise RuntimeError("Built-in Hugging Face session proxy URL is unavailable")
|
| 137 |
+
if direct_realtime_url:
|
| 138 |
+
logger.info("HF_REALTIME_CONNECTION_MODE=deployed; ignoring HF_REALTIME_WS_URL.")
|
| 139 |
+
|
| 140 |
+
allocator_headers = {"Authorization": f"Bearer {bearer_token}"} if bearer_token else None
|
| 141 |
+
async with httpx.AsyncClient(timeout=10.0) as http_client:
|
| 142 |
+
response = await http_client.post(session_url, headers=allocator_headers)
|
| 143 |
+
response.raise_for_status()
|
| 144 |
+
payload = response.json()
|
| 145 |
+
|
| 146 |
+
connect_url = payload.get("connect_url")
|
| 147 |
+
if not isinstance(connect_url, str) or not connect_url:
|
| 148 |
+
raise RuntimeError(f"Session allocator response did not contain a valid connect_url: {payload!r}")
|
| 149 |
+
|
| 150 |
+
parsed_connect_url = parse_hf_realtime_url(connect_url)
|
| 151 |
+
if not parsed_connect_url.has_realtime_path:
|
| 152 |
+
raise ValueError(f"Expected realtime connect URL ending with /realtime, got: {connect_url}")
|
| 153 |
+
|
| 154 |
+
logger.info("Allocated realtime session %s", payload.get("session_id") or "<unknown>")
|
| 155 |
+
client, connect_query = _build_openai_compatible_client_from_realtime_url(
|
| 156 |
+
connect_url,
|
| 157 |
+
bearer_token,
|
| 158 |
+
)
|
| 159 |
+
self._realtime_connect_query = connect_query
|
| 160 |
+
return client
|
src/reachy_mini_conversation_app/main.py
CHANGED
|
@@ -46,10 +46,25 @@ def run(
|
|
| 46 |
"""Run the Reachy Mini conversation app."""
|
| 47 |
# Putting these dependencies here makes the dashboard faster to load when the conversation app is installed
|
| 48 |
from reachy_mini_conversation_app.moves import MovementManager
|
| 49 |
-
from reachy_mini_conversation_app.config import
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 50 |
|
| 51 |
logger = setup_logger(args.debug)
|
| 52 |
logger.info("Starting Reachy Mini Conversation App")
|
|
|
|
| 53 |
|
| 54 |
if instance_path is not None:
|
| 55 |
try:
|
|
@@ -63,6 +78,26 @@ def run(
|
|
| 63 |
except Exception as e:
|
| 64 |
logger.warning("Failed to load instance configuration: %s", e)
|
| 65 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 66 |
from reachy_mini_conversation_app.console import LocalStream
|
| 67 |
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 68 |
from reachy_mini_conversation_app.audio.head_wobbler import HeadWobbler
|
|
@@ -144,40 +179,75 @@ def run(
|
|
| 144 |
if is_gemini_model():
|
| 145 |
from reachy_mini_conversation_app.gemini_live import GeminiLiveHandler
|
| 146 |
|
| 147 |
-
logger.info(
|
| 148 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 149 |
else:
|
| 150 |
from reachy_mini_conversation_app.openai_realtime import OpenaiRealtimeHandler
|
| 151 |
|
| 152 |
-
logger.info(
|
| 153 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 154 |
|
| 155 |
stream_manager: gr.Blocks | LocalStream | None = None
|
| 156 |
|
| 157 |
if args.gradio:
|
| 158 |
-
uses_gemini_backend = is_gemini_model()
|
| 159 |
-
api_key_textbox = gr.Textbox(
|
| 160 |
-
label="GEMINI_API_KEY" if uses_gemini_backend else "OPENAI API Key",
|
| 161 |
-
type="password",
|
| 162 |
-
value=(os.getenv("GEMINI_API_KEY") if uses_gemini_backend else os.getenv("OPENAI_API_KEY"))
|
| 163 |
-
if not get_space()
|
| 164 |
-
else "",
|
| 165 |
-
)
|
| 166 |
-
|
| 167 |
from reachy_mini_conversation_app.gradio_personality import PersonalityUI
|
| 168 |
|
| 169 |
personality_ui = PersonalityUI()
|
| 170 |
personality_ui.create_components()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 171 |
|
| 172 |
stream = Stream(
|
| 173 |
handler=handler,
|
| 174 |
mode="send-receive",
|
| 175 |
modality="audio",
|
| 176 |
-
additional_inputs=
|
| 177 |
-
chatbot,
|
| 178 |
-
api_key_textbox,
|
| 179 |
-
*personality_ui.additional_inputs_ordered(),
|
| 180 |
-
],
|
| 181 |
additional_outputs=[chatbot],
|
| 182 |
additional_outputs_handler=update_chatbot,
|
| 183 |
ui_args={"title": "Talk with Reachy Mini"},
|
|
|
|
| 46 |
"""Run the Reachy Mini conversation app."""
|
| 47 |
# Putting these dependencies here makes the dashboard faster to load when the conversation app is installed
|
| 48 |
from reachy_mini_conversation_app.moves import MovementManager
|
| 49 |
+
from reachy_mini_conversation_app.config import (
|
| 50 |
+
HF_BACKEND,
|
| 51 |
+
GEMINI_BACKEND,
|
| 52 |
+
OPENAI_BACKEND,
|
| 53 |
+
HF_LOCAL_CONNECTION_MODE,
|
| 54 |
+
config,
|
| 55 |
+
is_gemini_model,
|
| 56 |
+
get_backend_label,
|
| 57 |
+
get_hf_connection_selection,
|
| 58 |
+
refresh_runtime_config_from_env,
|
| 59 |
+
)
|
| 60 |
+
from reachy_mini_conversation_app.startup_settings import (
|
| 61 |
+
StartupSettings,
|
| 62 |
+
load_startup_settings_into_runtime,
|
| 63 |
+
)
|
| 64 |
|
| 65 |
logger = setup_logger(args.debug)
|
| 66 |
logger.info("Starting Reachy Mini Conversation App")
|
| 67 |
+
startup_settings = StartupSettings()
|
| 68 |
|
| 69 |
if instance_path is not None:
|
| 70 |
try:
|
|
|
|
| 78 |
except Exception as e:
|
| 79 |
logger.warning("Failed to load instance configuration: %s", e)
|
| 80 |
|
| 81 |
+
try:
|
| 82 |
+
startup_settings = load_startup_settings_into_runtime(instance_path)
|
| 83 |
+
except Exception as e:
|
| 84 |
+
logger.warning("Failed to load startup settings: %s", e)
|
| 85 |
+
|
| 86 |
+
if config.BACKEND_PROVIDER == HF_BACKEND:
|
| 87 |
+
logger.info(
|
| 88 |
+
"Configured backend provider: %s (%s), connection mode: %s",
|
| 89 |
+
config.BACKEND_PROVIDER,
|
| 90 |
+
get_backend_label(config.BACKEND_PROVIDER),
|
| 91 |
+
get_hf_connection_selection().mode,
|
| 92 |
+
)
|
| 93 |
+
else:
|
| 94 |
+
logger.info(
|
| 95 |
+
"Configured backend provider: %s (%s), model: %s",
|
| 96 |
+
config.BACKEND_PROVIDER,
|
| 97 |
+
get_backend_label(config.BACKEND_PROVIDER),
|
| 98 |
+
config.MODEL_NAME,
|
| 99 |
+
)
|
| 100 |
+
|
| 101 |
from reachy_mini_conversation_app.console import LocalStream
|
| 102 |
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 103 |
from reachy_mini_conversation_app.audio.head_wobbler import HeadWobbler
|
|
|
|
| 179 |
if is_gemini_model():
|
| 180 |
from reachy_mini_conversation_app.gemini_live import GeminiLiveHandler
|
| 181 |
|
| 182 |
+
logger.info(
|
| 183 |
+
"Using %s via GeminiLiveHandler",
|
| 184 |
+
get_backend_label(config.BACKEND_PROVIDER),
|
| 185 |
+
)
|
| 186 |
+
handler = GeminiLiveHandler(
|
| 187 |
+
deps,
|
| 188 |
+
gradio_mode=args.gradio,
|
| 189 |
+
instance_path=instance_path,
|
| 190 |
+
startup_voice=startup_settings.voice,
|
| 191 |
+
)
|
| 192 |
+
elif config.BACKEND_PROVIDER == HF_BACKEND:
|
| 193 |
+
from reachy_mini_conversation_app.huggingface_realtime import HuggingFaceRealtimeHandler
|
| 194 |
+
|
| 195 |
+
hf_connection_selection = get_hf_connection_selection()
|
| 196 |
+
transport_label = (
|
| 197 |
+
"Hugging Face direct websocket"
|
| 198 |
+
if hf_connection_selection.mode == HF_LOCAL_CONNECTION_MODE and hf_connection_selection.has_target
|
| 199 |
+
else "Hugging Face session proxy"
|
| 200 |
+
)
|
| 201 |
+
logger.info(
|
| 202 |
+
"Using %s via Hugging Face realtime handler (%s)",
|
| 203 |
+
get_backend_label(config.BACKEND_PROVIDER),
|
| 204 |
+
transport_label,
|
| 205 |
+
)
|
| 206 |
+
handler = HuggingFaceRealtimeHandler(
|
| 207 |
+
deps,
|
| 208 |
+
gradio_mode=args.gradio,
|
| 209 |
+
instance_path=instance_path,
|
| 210 |
+
startup_voice=startup_settings.voice,
|
| 211 |
+
) # type: ignore[assignment]
|
| 212 |
else:
|
| 213 |
from reachy_mini_conversation_app.openai_realtime import OpenaiRealtimeHandler
|
| 214 |
|
| 215 |
+
logger.info(
|
| 216 |
+
"Using %s via OpenAI realtime handler (OpenAI Realtime API)",
|
| 217 |
+
get_backend_label(config.BACKEND_PROVIDER),
|
| 218 |
+
)
|
| 219 |
+
handler = OpenaiRealtimeHandler(
|
| 220 |
+
deps,
|
| 221 |
+
gradio_mode=args.gradio,
|
| 222 |
+
instance_path=instance_path,
|
| 223 |
+
startup_voice=startup_settings.voice,
|
| 224 |
+
) # type: ignore[assignment]
|
| 225 |
|
| 226 |
stream_manager: gr.Blocks | LocalStream | None = None
|
| 227 |
|
| 228 |
if args.gradio:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 229 |
from reachy_mini_conversation_app.gradio_personality import PersonalityUI
|
| 230 |
|
| 231 |
personality_ui = PersonalityUI()
|
| 232 |
personality_ui.create_components()
|
| 233 |
+
additional_inputs: list[Any] = [chatbot, *personality_ui.additional_inputs_ordered()]
|
| 234 |
+
|
| 235 |
+
if config.BACKEND_PROVIDER in {OPENAI_BACKEND, GEMINI_BACKEND}:
|
| 236 |
+
uses_gemini_backend = is_gemini_model()
|
| 237 |
+
api_key_textbox = gr.Textbox(
|
| 238 |
+
label="GEMINI_API_KEY" if uses_gemini_backend else "OPENAI API Key",
|
| 239 |
+
type="password",
|
| 240 |
+
value=(os.getenv("GEMINI_API_KEY") if uses_gemini_backend else os.getenv("OPENAI_API_KEY"))
|
| 241 |
+
if not get_space()
|
| 242 |
+
else "",
|
| 243 |
+
)
|
| 244 |
+
additional_inputs.insert(1, api_key_textbox)
|
| 245 |
|
| 246 |
stream = Stream(
|
| 247 |
handler=handler,
|
| 248 |
mode="send-receive",
|
| 249 |
modality="audio",
|
| 250 |
+
additional_inputs=additional_inputs,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 251 |
additional_outputs=[chatbot],
|
| 252 |
additional_outputs_handler=update_chatbot,
|
| 253 |
ui_args={"title": "Talk with Reachy Mini"},
|
src/reachy_mini_conversation_app/openai_realtime.py
CHANGED
|
@@ -1,830 +1,167 @@
|
|
| 1 |
-
import json
|
| 2 |
-
import uuid
|
| 3 |
-
import base64
|
| 4 |
-
import random
|
| 5 |
-
import asyncio
|
| 6 |
import logging
|
| 7 |
-
from typing import Any,
|
| 8 |
from pathlib import Path
|
| 9 |
-
from datetime import datetime
|
| 10 |
|
| 11 |
-
import numpy as np
|
| 12 |
-
import gradio as gr
|
| 13 |
from openai import AsyncOpenAI
|
| 14 |
-
from fastrtc import AdditionalOutputs, AsyncStreamHandler, wait_for_item, audio_to_int16
|
| 15 |
-
from pydantic import Field, BaseModel
|
| 16 |
-
from numpy.typing import NDArray
|
| 17 |
-
from scipy.signal import resample
|
| 18 |
from openai.types.realtime import (
|
| 19 |
AudioTranscriptionParam,
|
| 20 |
RealtimeAudioConfigParam,
|
| 21 |
RealtimeAudioConfigInputParam,
|
| 22 |
RealtimeAudioConfigOutputParam,
|
| 23 |
-
RealtimeResponseCreateParamsParam,
|
| 24 |
RealtimeSessionCreateRequestParam,
|
| 25 |
)
|
| 26 |
-
from websockets.exceptions import ConnectionClosedError
|
| 27 |
-
from openai.resources.realtime.realtime import AsyncRealtimeConnection
|
| 28 |
from openai.types.realtime.realtime_audio_formats_param import AudioPCM
|
| 29 |
from openai.types.realtime.realtime_audio_input_turn_detection_param import ServerVad
|
| 30 |
|
| 31 |
-
from reachy_mini_conversation_app.config import
|
| 32 |
from reachy_mini_conversation_app.prompts import get_session_voice, get_session_instructions
|
| 33 |
-
from reachy_mini_conversation_app.
|
| 34 |
-
|
| 35 |
-
get_tool_specs,
|
| 36 |
-
)
|
| 37 |
-
from reachy_mini_conversation_app.tools.background_tool_manager import (
|
| 38 |
-
ToolCallRoutine,
|
| 39 |
-
ToolNotification,
|
| 40 |
-
BackgroundToolManager,
|
| 41 |
-
)
|
| 42 |
|
| 43 |
|
| 44 |
logger = logging.getLogger(__name__)
|
| 45 |
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
cost = 0.0
|
| 71 |
-
if inp:
|
| 72 |
-
cost += (getattr(inp, "audio_tokens", 0) or 0) * AUDIO_INPUT_COST_PER_1M / 1e6
|
| 73 |
-
cost += (getattr(inp, "text_tokens", 0) or 0) * TEXT_INPUT_COST_PER_1M / 1e6
|
| 74 |
-
cost += (getattr(inp, "image_tokens", 0) or 0) * IMAGE_INPUT_COST_PER_1M / 1e6
|
| 75 |
-
if out:
|
| 76 |
-
cost += (getattr(out, "audio_tokens", 0) or 0) * AUDIO_OUTPUT_COST_PER_1M / 1e6
|
| 77 |
-
cost += (getattr(out, "text_tokens", 0) or 0) * TEXT_OUTPUT_COST_PER_1M / 1e6
|
| 78 |
-
return cost
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
class OpenaiRealtimeHandler(AsyncStreamHandler):
|
| 82 |
-
"""An OpenAI realtime handler for fastrtc Stream."""
|
| 83 |
-
|
| 84 |
-
def __init__(self, deps: ToolDependencies, gradio_mode: bool = False, instance_path: Optional[str] = None):
|
| 85 |
-
"""Initialize the handler."""
|
| 86 |
-
super().__init__(
|
| 87 |
-
expected_layout="mono",
|
| 88 |
-
output_sample_rate=OPEN_AI_OUTPUT_SAMPLE_RATE,
|
| 89 |
-
input_sample_rate=OPEN_AI_INPUT_SAMPLE_RATE,
|
| 90 |
-
)
|
| 91 |
-
|
| 92 |
-
# Override typing of the sample rates to match OpenAI's requirements
|
| 93 |
-
self.output_sample_rate: Literal[24000] = self.output_sample_rate
|
| 94 |
-
self.input_sample_rate: Literal[24000] = self.input_sample_rate
|
| 95 |
-
|
| 96 |
-
self.deps = deps
|
| 97 |
-
|
| 98 |
-
# Override type annotations for OpenAI strict typing (only for values used in API)
|
| 99 |
-
self.output_sample_rate = OPEN_AI_OUTPUT_SAMPLE_RATE
|
| 100 |
-
self.input_sample_rate = OPEN_AI_INPUT_SAMPLE_RATE
|
| 101 |
-
|
| 102 |
-
self.connection: AsyncRealtimeConnection | None = None
|
| 103 |
-
self.output_queue: "asyncio.Queue[Tuple[int, NDArray[np.int16]] | AdditionalOutputs]" = asyncio.Queue()
|
| 104 |
-
|
| 105 |
-
self.last_activity_time = asyncio.get_event_loop().time()
|
| 106 |
-
self.start_time = asyncio.get_event_loop().time()
|
| 107 |
-
self.is_idle_tool_call = False
|
| 108 |
-
self.gradio_mode = gradio_mode
|
| 109 |
-
self.instance_path = instance_path
|
| 110 |
-
self._voice_override: str | None = None
|
| 111 |
-
# Track how the API key was provided (env vs textbox) and its value
|
| 112 |
self._key_source: Literal["env", "textbox"] = "env"
|
| 113 |
self._provided_api_key: str | None = None
|
| 114 |
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
self.partial_debounce_delay = 0.5 # seconds
|
| 118 |
-
self.input_transcript_chunks_by_item = InputTranscriptChunksByItem()
|
| 119 |
-
|
| 120 |
-
# Internal lifecycle flags
|
| 121 |
-
self._connected_event: asyncio.Event = asyncio.Event()
|
| 122 |
-
|
| 123 |
-
# Background tool manager
|
| 124 |
-
self.tool_manager = BackgroundToolManager()
|
| 125 |
-
|
| 126 |
-
# Cost tracking
|
| 127 |
-
self.cumulative_cost: float = 0.0
|
| 128 |
-
|
| 129 |
-
# Response-in-progress guard: the Realtime API only allows one active
|
| 130 |
-
# response per conversation at a time. A dedicated worker task
|
| 131 |
-
# (_response_sender_loop) dequeues and sends one request at a time
|
| 132 |
-
self._pending_responses: asyncio.Queue[dict[str, Any]] = asyncio.Queue()
|
| 133 |
-
self._response_done_event: asyncio.Event = asyncio.Event()
|
| 134 |
-
self._response_done_event.set()
|
| 135 |
-
self._last_response_rejected: bool = False
|
| 136 |
-
|
| 137 |
-
def copy(self) -> "OpenaiRealtimeHandler":
|
| 138 |
-
"""Create a copy of the handler."""
|
| 139 |
-
return OpenaiRealtimeHandler(self.deps, self.gradio_mode, self.instance_path)
|
| 140 |
-
|
| 141 |
-
async def apply_personality(self, profile: str | None) -> str:
|
| 142 |
-
"""Apply a new personality (profile) at runtime if possible.
|
| 143 |
-
|
| 144 |
-
- Updates the global config's selected profile for subsequent calls.
|
| 145 |
-
- If a realtime connection is active, sends a session.update with the
|
| 146 |
-
freshly resolved instructions so the change takes effect immediately.
|
| 147 |
-
|
| 148 |
-
Returns a short status message for UI feedback.
|
| 149 |
-
"""
|
| 150 |
-
try:
|
| 151 |
-
# Update the in-process config value and env
|
| 152 |
-
from reachy_mini_conversation_app.config import config as _config
|
| 153 |
-
from reachy_mini_conversation_app.config import set_custom_profile
|
| 154 |
-
|
| 155 |
-
set_custom_profile(profile)
|
| 156 |
-
self._voice_override = None
|
| 157 |
-
logger.info(
|
| 158 |
-
"Set custom profile to %r (config=%r)", profile, getattr(_config, "REACHY_MINI_CUSTOM_PROFILE", None)
|
| 159 |
-
)
|
| 160 |
-
|
| 161 |
-
try:
|
| 162 |
-
instructions = get_session_instructions()
|
| 163 |
-
voice = get_session_voice()
|
| 164 |
-
except BaseException as e: # catch SystemExit from prompt loader without crashing
|
| 165 |
-
logger.error("Failed to resolve personality content: %s", e)
|
| 166 |
-
return f"Failed to apply personality: {e}"
|
| 167 |
-
|
| 168 |
-
# Attempt a live update first, then force a full restart to ensure it sticks
|
| 169 |
-
if self.connection is not None:
|
| 170 |
-
try:
|
| 171 |
-
await self.connection.session.update(
|
| 172 |
-
session=RealtimeSessionCreateRequestParam(
|
| 173 |
-
type="realtime",
|
| 174 |
-
instructions=instructions,
|
| 175 |
-
audio=RealtimeAudioConfigParam(
|
| 176 |
-
output=RealtimeAudioConfigOutputParam(
|
| 177 |
-
voice=voice,
|
| 178 |
-
),
|
| 179 |
-
),
|
| 180 |
-
),
|
| 181 |
-
)
|
| 182 |
-
logger.info("Applied personality via live update: %s", profile or "built-in default")
|
| 183 |
-
except Exception as e:
|
| 184 |
-
logger.warning("Live update failed; will restart session: %s", e)
|
| 185 |
-
|
| 186 |
-
# Force a real restart to guarantee the new instructions/voice
|
| 187 |
-
try:
|
| 188 |
-
await self._restart_session()
|
| 189 |
-
return "Applied personality and restarted realtime session."
|
| 190 |
-
except Exception as e:
|
| 191 |
-
logger.warning("Failed to restart session after apply: %s", e)
|
| 192 |
-
return "Applied personality. Will take effect on next connection."
|
| 193 |
-
else:
|
| 194 |
-
logger.info(
|
| 195 |
-
"Applied personality recorded: %s (no live connection; will apply on next session)",
|
| 196 |
-
profile or "built-in default",
|
| 197 |
-
)
|
| 198 |
-
return "Applied personality. Will take effect on next connection."
|
| 199 |
-
except Exception as e:
|
| 200 |
-
logger.error("Error applying personality '%s': %s", profile, e)
|
| 201 |
-
return f"Failed to apply personality: {e}"
|
| 202 |
-
|
| 203 |
-
async def change_voice(self, voice: str) -> str:
|
| 204 |
-
"""Change only the voice and restart the session."""
|
| 205 |
-
self._voice_override = voice
|
| 206 |
-
if getattr(self, "client", None) is not None:
|
| 207 |
-
try:
|
| 208 |
-
await self._restart_session()
|
| 209 |
-
return f"Voice changed to {voice}."
|
| 210 |
-
except Exception as e:
|
| 211 |
-
logger.warning("Failed to restart session for voice change: %s", e)
|
| 212 |
-
return "Voice change failed. Will take effect on next connection."
|
| 213 |
-
return "Voice changed. Will take effect on next connection."
|
| 214 |
-
|
| 215 |
-
def get_current_voice(self) -> str:
|
| 216 |
-
"""Return the voice currently selected for this handler."""
|
| 217 |
-
return self._voice_override or get_session_voice()
|
| 218 |
-
|
| 219 |
-
async def _emit_debounced_partial(self, transcript: str, item_id: str, sequence_counter: int) -> None:
|
| 220 |
-
"""Emit partial transcript after debounce delay."""
|
| 221 |
-
try:
|
| 222 |
-
await asyncio.sleep(self.partial_debounce_delay)
|
| 223 |
-
|
| 224 |
-
input_transcript = self.input_transcript_chunks_by_item
|
| 225 |
-
if input_transcript.item_id == item_id and len(input_transcript.deltas) - 1 == sequence_counter:
|
| 226 |
-
await self.output_queue.put(AdditionalOutputs({"role": "user_partial", "content": transcript}))
|
| 227 |
-
logger.debug(f"Debounced partial emitted: {transcript}")
|
| 228 |
-
except asyncio.CancelledError:
|
| 229 |
-
logger.debug("Debounced partial cancelled")
|
| 230 |
-
raise
|
| 231 |
-
|
| 232 |
-
async def start_up(self) -> None:
|
| 233 |
-
"""Start the handler with minimal retries on unexpected websocket closure."""
|
| 234 |
openai_api_key = config.OPENAI_API_KEY
|
| 235 |
-
if self.gradio_mode
|
| 236 |
-
|
| 237 |
-
await self.wait_for_args() # type: ignore[no-untyped-call]
|
| 238 |
-
args = list(self.latest_args)
|
| 239 |
-
textbox_api_key = args[3] if len(args[3]) > 0 else None
|
| 240 |
-
if textbox_api_key is not None:
|
| 241 |
-
openai_api_key = textbox_api_key
|
| 242 |
-
self._key_source = "textbox"
|
| 243 |
-
self._provided_api_key = textbox_api_key
|
| 244 |
-
else:
|
| 245 |
-
openai_api_key = config.OPENAI_API_KEY
|
| 246 |
-
else:
|
| 247 |
-
if not openai_api_key or not openai_api_key.strip():
|
| 248 |
-
# In headless console mode, LocalStream now blocks startup until the key is provided.
|
| 249 |
-
# However, unit tests may invoke this handler directly with a stubbed client.
|
| 250 |
-
# To keep tests hermetic without requiring a real key, fall back to a placeholder.
|
| 251 |
-
logger.warning("OPENAI_API_KEY missing. Proceeding with a placeholder (tests/offline).")
|
| 252 |
-
openai_api_key = "DUMMY"
|
| 253 |
|
| 254 |
-
self.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 255 |
|
| 256 |
-
|
| 257 |
-
|
| 258 |
-
|
| 259 |
-
|
| 260 |
-
|
| 261 |
return
|
| 262 |
-
except ConnectionClosedError as e:
|
| 263 |
-
# Abrupt close (e.g., "no close frame received or sent") → retry
|
| 264 |
-
logger.warning("Realtime websocket closed unexpectedly (attempt %d/%d): %s", attempt, max_attempts, e)
|
| 265 |
-
if attempt < max_attempts:
|
| 266 |
-
# exponential backoff with jitter
|
| 267 |
-
base_delay = 2 ** (attempt - 1) # 1s, 2s, 4s, 8s, etc.
|
| 268 |
-
jitter = random.uniform(0, 0.5)
|
| 269 |
-
delay = base_delay + jitter
|
| 270 |
-
logger.info("Retrying in %.1f seconds...", delay)
|
| 271 |
-
await asyncio.sleep(delay)
|
| 272 |
-
continue
|
| 273 |
-
raise
|
| 274 |
-
finally:
|
| 275 |
-
# never keep a stale reference
|
| 276 |
-
self.connection = None
|
| 277 |
-
try:
|
| 278 |
-
self._connected_event.clear()
|
| 279 |
-
except Exception:
|
| 280 |
-
pass
|
| 281 |
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
Does not block the caller while the new session is establishing.
|
| 286 |
-
"""
|
| 287 |
-
try:
|
| 288 |
-
if self.connection is not None:
|
| 289 |
-
try:
|
| 290 |
-
await self.connection.close()
|
| 291 |
-
except Exception:
|
| 292 |
-
pass
|
| 293 |
-
finally:
|
| 294 |
-
self.connection = None
|
| 295 |
|
| 296 |
-
|
| 297 |
-
if
|
| 298 |
-
logger.warning("
|
|
|
|
|
|
|
|
|
|
| 299 |
return
|
| 300 |
|
| 301 |
-
#
|
| 302 |
-
try:
|
| 303 |
-
self._connected_event.clear()
|
| 304 |
-
except Exception:
|
| 305 |
-
pass
|
| 306 |
-
asyncio.create_task(self._run_realtime_session(), name="openai-realtime-restart")
|
| 307 |
try:
|
| 308 |
-
|
| 309 |
-
logger.info("Realtime session restarted and connected.")
|
| 310 |
-
except asyncio.TimeoutError:
|
| 311 |
-
logger.warning("Realtime session restart timed out; continuing in background.")
|
| 312 |
-
except Exception as e:
|
| 313 |
-
logger.warning("_restart_session failed: %s", e)
|
| 314 |
-
|
| 315 |
-
async def _safe_response_create(self, **kwargs: Any) -> None:
|
| 316 |
-
"""Enqueue a response.create() kwargs for the sender worker _response_sender_loop().
|
| 317 |
-
|
| 318 |
-
This method never blocks the caller.
|
| 319 |
-
"""
|
| 320 |
-
await self._pending_responses.put(kwargs)
|
| 321 |
-
|
| 322 |
-
async def _response_sender_loop(self) -> None:
|
| 323 |
-
"""Dedicated worker that sends ``response.create()`` calls serially.
|
| 324 |
|
| 325 |
-
|
| 326 |
-
|
|
|
|
| 327 |
|
| 328 |
-
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
|
| 332 |
-
|
| 333 |
-
"""
|
| 334 |
-
while self.connection:
|
| 335 |
-
try:
|
| 336 |
-
kwargs = await self._pending_responses.get()
|
| 337 |
-
except asyncio.CancelledError:
|
| 338 |
return
|
| 339 |
|
| 340 |
-
|
| 341 |
-
|
| 342 |
-
|
| 343 |
-
while not sent and self.connection and attempts < max_retries:
|
| 344 |
-
try:
|
| 345 |
-
await asyncio.wait_for(self._response_done_event.wait(), timeout=_RESPONSE_DONE_TIMEOUT)
|
| 346 |
-
except asyncio.TimeoutError:
|
| 347 |
-
logger.debug("Timed out waiting for previous response to finish; forcing ahead")
|
| 348 |
-
self._response_done_event.set()
|
| 349 |
-
|
| 350 |
-
if not self.connection:
|
| 351 |
-
break
|
| 352 |
-
|
| 353 |
-
self._last_response_rejected = False
|
| 354 |
try:
|
| 355 |
-
|
|
|
|
| 356 |
except Exception as e:
|
| 357 |
-
logger.
|
| 358 |
-
self._response_done_event.set()
|
| 359 |
-
break
|
| 360 |
|
| 361 |
-
|
| 362 |
-
|
| 363 |
-
|
| 364 |
-
|
| 365 |
-
|
| 366 |
break
|
|
|
|
|
|
|
| 367 |
|
| 368 |
-
|
| 369 |
-
|
| 370 |
-
|
| 371 |
-
if attempts >= max_retries:
|
| 372 |
-
logger.debug("response.create rejected %d times; giving up", attempts)
|
| 373 |
-
break
|
| 374 |
-
logger.debug("response.create was rejected; retrying (%d/%d)", attempts, max_retries)
|
| 375 |
-
continue
|
| 376 |
-
|
| 377 |
-
sent = True
|
| 378 |
-
|
| 379 |
-
async def _handle_tool_result(self, bg_tool: ToolNotification) -> None:
|
| 380 |
-
"""Process the result of a tool call."""
|
| 381 |
-
if bg_tool.error is not None:
|
| 382 |
-
logger.error("Tool '%s' (id=%s) failed with error: %s", bg_tool.tool_name, bg_tool.id, bg_tool.error)
|
| 383 |
-
tool_result = {"error": bg_tool.error}
|
| 384 |
-
elif bg_tool.result is not None:
|
| 385 |
-
tool_result = bg_tool.result
|
| 386 |
-
logger.info(
|
| 387 |
-
"Tool '%s' (id=%s) executed successfully.",
|
| 388 |
-
bg_tool.tool_name,
|
| 389 |
-
bg_tool.id,
|
| 390 |
-
)
|
| 391 |
-
logger.debug("Tool '%s' full result: %s", bg_tool.tool_name, tool_result)
|
| 392 |
-
else:
|
| 393 |
-
logger.warning("Tool '%s' (id=%s) returned no result and no error", bg_tool.tool_name, bg_tool.id)
|
| 394 |
-
tool_result = {"error": "No result returned from tool execution"}
|
| 395 |
-
|
| 396 |
-
# Connection may have closed while tool was running
|
| 397 |
-
if not self.connection:
|
| 398 |
-
logger.warning(
|
| 399 |
-
"Connection closed during tool '%s' (id=%s) execution; cannot send result back",
|
| 400 |
-
bg_tool.tool_name,
|
| 401 |
-
bg_tool.id,
|
| 402 |
-
)
|
| 403 |
-
return
|
| 404 |
-
|
| 405 |
-
try:
|
| 406 |
-
# TODO: refactor this since it's repeated here, in the camera branch below, and in send_idle_signal
|
| 407 |
-
if isinstance(bg_tool.id, str):
|
| 408 |
-
await self.connection.conversation.item.create(
|
| 409 |
-
item={
|
| 410 |
-
"type": "function_call_output",
|
| 411 |
-
"call_id": bg_tool.id,
|
| 412 |
-
"output": json.dumps(tool_result),
|
| 413 |
-
},
|
| 414 |
-
)
|
| 415 |
-
|
| 416 |
-
await self.output_queue.put(
|
| 417 |
-
AdditionalOutputs(
|
| 418 |
-
{
|
| 419 |
-
"role": "assistant",
|
| 420 |
-
"content": json.dumps(tool_result),
|
| 421 |
-
# Gradio UI metadata.status accept only "pending" and "done". Do not accept bg.tool.status values.
|
| 422 |
-
"metadata": {
|
| 423 |
-
"title": f"🛠️ Used tool {bg_tool.tool_name}",
|
| 424 |
-
"status": "done",
|
| 425 |
-
},
|
| 426 |
-
},
|
| 427 |
-
),
|
| 428 |
-
)
|
| 429 |
-
|
| 430 |
-
if bg_tool.tool_name == "camera" and "b64_im" in tool_result:
|
| 431 |
-
# use raw base64, don't json.dumps (which adds quotes)
|
| 432 |
-
b64_im = tool_result["b64_im"]
|
| 433 |
-
if not isinstance(b64_im, str):
|
| 434 |
-
logger.warning("Unexpected type for b64_im: %s", type(b64_im))
|
| 435 |
-
b64_im = str(b64_im)
|
| 436 |
-
await self.connection.conversation.item.create(
|
| 437 |
-
item={
|
| 438 |
-
"type": "message",
|
| 439 |
-
"role": "user",
|
| 440 |
-
"content": [
|
| 441 |
-
{
|
| 442 |
-
"type": "input_image",
|
| 443 |
-
"image_url": f"data:image/jpeg;base64,{b64_im}",
|
| 444 |
-
},
|
| 445 |
-
],
|
| 446 |
-
},
|
| 447 |
-
)
|
| 448 |
-
logger.info("Added camera image to conversation")
|
| 449 |
-
|
| 450 |
-
if self.deps.camera_worker is not None:
|
| 451 |
-
np_img = self.deps.camera_worker.get_latest_frame()
|
| 452 |
-
if np_img is not None:
|
| 453 |
-
rgb_frame = np.ascontiguousarray(np_img[..., ::-1])
|
| 454 |
-
else:
|
| 455 |
-
rgb_frame = None
|
| 456 |
-
img = gr.Image(value=rgb_frame)
|
| 457 |
-
|
| 458 |
-
await self.output_queue.put(
|
| 459 |
-
AdditionalOutputs(
|
| 460 |
-
{
|
| 461 |
-
"role": "assistant",
|
| 462 |
-
"content": img,
|
| 463 |
-
},
|
| 464 |
-
),
|
| 465 |
-
)
|
| 466 |
-
|
| 467 |
-
# If this tool call was triggered by an idle signal, don't make the robot speak.
|
| 468 |
-
# For other tool calls, let the robot reply out loud.
|
| 469 |
-
if not bg_tool.is_idle_tool_call:
|
| 470 |
-
await self._safe_response_create(
|
| 471 |
-
response=RealtimeResponseCreateParamsParam(
|
| 472 |
-
instructions="Use the tool result just returned and answer concisely in speech.",
|
| 473 |
-
),
|
| 474 |
-
)
|
| 475 |
-
|
| 476 |
-
except ConnectionClosedError:
|
| 477 |
-
logger.warning("Connection closed while sending tool result")
|
| 478 |
-
self.connection = None
|
| 479 |
-
self._response_done_event.set()
|
| 480 |
-
|
| 481 |
-
async def _run_realtime_session(self) -> None:
|
| 482 |
-
"""Establish and manage a single realtime session."""
|
| 483 |
-
async with self.client.realtime.connect(model=config.MODEL_NAME) as conn:
|
| 484 |
-
try:
|
| 485 |
-
session_config = RealtimeSessionCreateRequestParam(
|
| 486 |
-
type="realtime",
|
| 487 |
-
instructions=get_session_instructions(),
|
| 488 |
-
audio=RealtimeAudioConfigParam(
|
| 489 |
-
input=RealtimeAudioConfigInputParam(
|
| 490 |
-
format=AudioPCM(type="audio/pcm", rate=self.input_sample_rate),
|
| 491 |
-
transcription=AudioTranscriptionParam(model="gpt-4o-transcribe", language="en"),
|
| 492 |
-
turn_detection=ServerVad(type="server_vad", interrupt_response=True),
|
| 493 |
-
),
|
| 494 |
-
output=RealtimeAudioConfigOutputParam(
|
| 495 |
-
format=AudioPCM(type="audio/pcm", rate=self.output_sample_rate),
|
| 496 |
-
voice=self._voice_override or get_session_voice(),
|
| 497 |
-
),
|
| 498 |
-
),
|
| 499 |
-
tools=get_tool_specs(), # type: ignore[typeddict-item]
|
| 500 |
-
tool_choice="auto",
|
| 501 |
-
)
|
| 502 |
-
await conn.session.update(session=session_config)
|
| 503 |
-
logger.info(
|
| 504 |
-
"Realtime session initialized with profile=%r voice=%r",
|
| 505 |
-
getattr(config, "REACHY_MINI_CUSTOM_PROFILE", None),
|
| 506 |
-
self._voice_override or get_session_voice(),
|
| 507 |
-
)
|
| 508 |
-
# If we reached here, the session update succeeded which implies the API key worked.
|
| 509 |
-
# Persist the key to a newly created .env (copied from .env.example) if needed.
|
| 510 |
-
self._persist_api_key_if_needed()
|
| 511 |
-
except Exception:
|
| 512 |
-
logger.exception("Realtime session.update failed; aborting startup")
|
| 513 |
-
return
|
| 514 |
-
|
| 515 |
-
logger.info("Realtime session updated successfully")
|
| 516 |
-
|
| 517 |
-
# Reset the partial-transcript accumulator for each new session
|
| 518 |
-
self.input_transcript_chunks_by_item = InputTranscriptChunksByItem()
|
| 519 |
-
|
| 520 |
-
# Manage event received from the openai server
|
| 521 |
-
self.connection = conn
|
| 522 |
-
try:
|
| 523 |
-
self._connected_event.set()
|
| 524 |
-
except Exception:
|
| 525 |
-
pass
|
| 526 |
-
|
| 527 |
-
response_sender_task: asyncio.Task[None] | None = None
|
| 528 |
-
try:
|
| 529 |
-
# Start the background tool manager
|
| 530 |
-
self.tool_manager.start_up(tool_callbacks=[self._handle_tool_result])
|
| 531 |
-
|
| 532 |
-
# Start the response sender worker
|
| 533 |
-
response_sender_task = asyncio.create_task(self._response_sender_loop(), name="response-sender")
|
| 534 |
-
|
| 535 |
-
async for event in self.connection:
|
| 536 |
-
logger.debug(f"OpenAI event: {event.type}")
|
| 537 |
-
if event.type == "input_audio_buffer.speech_started":
|
| 538 |
-
if hasattr(self, "_clear_queue") and callable(self._clear_queue):
|
| 539 |
-
self._clear_queue()
|
| 540 |
-
if self.deps.head_wobbler is not None:
|
| 541 |
-
self.deps.head_wobbler.reset()
|
| 542 |
-
self.deps.movement_manager.set_listening(True)
|
| 543 |
-
logger.debug("User speech started")
|
| 544 |
-
|
| 545 |
-
if event.type == "input_audio_buffer.speech_stopped":
|
| 546 |
-
self.deps.movement_manager.set_listening(False)
|
| 547 |
-
logger.debug("User speech stopped - server will auto-commit with VAD")
|
| 548 |
-
|
| 549 |
-
if event.type == "response.output_audio.done":
|
| 550 |
-
if self.deps.head_wobbler is not None:
|
| 551 |
-
self.deps.head_wobbler.request_reset_after_current_audio()
|
| 552 |
-
logger.debug("response completed")
|
| 553 |
-
|
| 554 |
-
if event.type == "response.created":
|
| 555 |
-
self._response_done_event.clear()
|
| 556 |
-
logger.debug("Response created (active)")
|
| 557 |
-
|
| 558 |
-
if event.type == "response.done":
|
| 559 |
-
# Doesn't mean the audio is done playing
|
| 560 |
-
self._response_done_event.set()
|
| 561 |
-
logger.debug("Response done")
|
| 562 |
-
|
| 563 |
-
response = getattr(event, "response", None)
|
| 564 |
-
usage = getattr(response, "usage", None) if response else None
|
| 565 |
-
if usage:
|
| 566 |
-
cost = _compute_response_cost(usage)
|
| 567 |
-
self.cumulative_cost += cost
|
| 568 |
-
logger.debug("Cost: $%.4f | Cumulative: $%.4f", cost, self.cumulative_cost)
|
| 569 |
-
else:
|
| 570 |
-
logger.warning("No usage data available for cost tracking")
|
| 571 |
-
|
| 572 |
-
if event.type == "conversation.item.input_audio_transcription.delta":
|
| 573 |
-
logger.debug(f"User partial transcript: {event.delta}")
|
| 574 |
-
|
| 575 |
-
item_id = event.item_id
|
| 576 |
-
delta = event.delta or ""
|
| 577 |
-
|
| 578 |
-
input_transcript = self.input_transcript_chunks_by_item
|
| 579 |
-
if input_transcript.item_id != item_id:
|
| 580 |
-
input_transcript.item_id = item_id
|
| 581 |
-
input_transcript.deltas = [delta]
|
| 582 |
-
else:
|
| 583 |
-
input_transcript.deltas.append(delta)
|
| 584 |
-
|
| 585 |
-
sequence_counter = len(input_transcript.deltas) - 1
|
| 586 |
-
|
| 587 |
-
# Cancel previous debounce task if it exists
|
| 588 |
-
if self.partial_transcript_task and not self.partial_transcript_task.done():
|
| 589 |
-
self.partial_transcript_task.cancel()
|
| 590 |
-
try:
|
| 591 |
-
await self.partial_transcript_task
|
| 592 |
-
except asyncio.CancelledError:
|
| 593 |
-
pass
|
| 594 |
-
|
| 595 |
-
# Start new debounce timer with the last delta
|
| 596 |
-
self.partial_transcript_task = asyncio.create_task(
|
| 597 |
-
self._emit_debounced_partial("".join(input_transcript.deltas), item_id, sequence_counter)
|
| 598 |
-
)
|
| 599 |
-
|
| 600 |
-
# Handle completed transcription (user finished speaking)
|
| 601 |
-
if event.type == "conversation.item.input_audio_transcription.completed":
|
| 602 |
-
logger.debug(f"User transcript: {event.transcript}")
|
| 603 |
-
|
| 604 |
-
# Cancel any pending partial emission
|
| 605 |
-
if self.partial_transcript_task and not self.partial_transcript_task.done():
|
| 606 |
-
self.partial_transcript_task.cancel()
|
| 607 |
-
try:
|
| 608 |
-
await self.partial_transcript_task
|
| 609 |
-
except asyncio.CancelledError:
|
| 610 |
-
pass
|
| 611 |
-
|
| 612 |
-
await self.output_queue.put(AdditionalOutputs({"role": "user", "content": event.transcript}))
|
| 613 |
-
|
| 614 |
-
# Handle assistant transcription
|
| 615 |
-
if event.type == "response.output_audio_transcript.done":
|
| 616 |
-
logger.debug(f"Assistant transcript: {event.transcript}")
|
| 617 |
-
await self.output_queue.put(
|
| 618 |
-
AdditionalOutputs({"role": "assistant", "content": event.transcript})
|
| 619 |
-
)
|
| 620 |
-
|
| 621 |
-
# Handle audio delta
|
| 622 |
-
if event.type == "response.output_audio.delta":
|
| 623 |
-
if self.gradio_mode and self.deps.head_wobbler is not None:
|
| 624 |
-
self.deps.head_wobbler.feed(event.delta)
|
| 625 |
-
self.last_activity_time = asyncio.get_event_loop().time()
|
| 626 |
-
logger.debug("last activity time updated to %s", self.last_activity_time)
|
| 627 |
-
await self.output_queue.put(
|
| 628 |
-
(
|
| 629 |
-
self.output_sample_rate,
|
| 630 |
-
np.frombuffer(base64.b64decode(event.delta), dtype=np.int16).reshape(1, -1),
|
| 631 |
-
),
|
| 632 |
-
)
|
| 633 |
-
|
| 634 |
-
# ---- tool-calling plumbing ----
|
| 635 |
-
if event.type == "response.function_call_arguments.done":
|
| 636 |
-
tool_name = getattr(event, "name", None)
|
| 637 |
-
args_json_str = getattr(event, "arguments", None)
|
| 638 |
-
call_id: str = str(getattr(event, "call_id", uuid.uuid4()))
|
| 639 |
-
|
| 640 |
-
logger.info(
|
| 641 |
-
"Tool call received — tool_name=%r, call_id=%s, is_idle=%s, args=%s",
|
| 642 |
-
tool_name,
|
| 643 |
-
call_id,
|
| 644 |
-
self.is_idle_tool_call,
|
| 645 |
-
args_json_str,
|
| 646 |
-
)
|
| 647 |
-
|
| 648 |
-
if not isinstance(tool_name, str) or not isinstance(args_json_str, str):
|
| 649 |
-
logger.error(
|
| 650 |
-
"Invalid tool call: tool_name=%s (type=%s), args=%s (type=%s), call_id=%s",
|
| 651 |
-
tool_name,
|
| 652 |
-
type(tool_name).__name__,
|
| 653 |
-
args_json_str,
|
| 654 |
-
type(args_json_str).__name__,
|
| 655 |
-
call_id,
|
| 656 |
-
)
|
| 657 |
-
continue
|
| 658 |
-
|
| 659 |
-
bg_tool = await self.tool_manager.start_tool(
|
| 660 |
-
call_id=call_id,
|
| 661 |
-
tool_call_routine=ToolCallRoutine(
|
| 662 |
-
tool_name=tool_name,
|
| 663 |
-
args_json_str=args_json_str,
|
| 664 |
-
deps=self.deps,
|
| 665 |
-
),
|
| 666 |
-
is_idle_tool_call=self.is_idle_tool_call,
|
| 667 |
-
)
|
| 668 |
-
|
| 669 |
-
await self.output_queue.put(
|
| 670 |
-
AdditionalOutputs(
|
| 671 |
-
{
|
| 672 |
-
"role": "assistant",
|
| 673 |
-
"content": f"🛠️ Used tool {tool_name} with args {args_json_str}. The tool is now running. Tool ID: {bg_tool.tool_id}",
|
| 674 |
-
},
|
| 675 |
-
),
|
| 676 |
-
)
|
| 677 |
-
|
| 678 |
-
if self.is_idle_tool_call:
|
| 679 |
-
self.is_idle_tool_call = False
|
| 680 |
-
|
| 681 |
-
logger.info(
|
| 682 |
-
"Started background tool: %s (id=%s, call_id=%s)", tool_name, bg_tool.tool_id, call_id
|
| 683 |
-
)
|
| 684 |
-
|
| 685 |
-
# server error
|
| 686 |
-
if event.type == "error":
|
| 687 |
-
err = getattr(event, "error", None)
|
| 688 |
-
msg = getattr(err, "message", str(err) if err else "unknown error")
|
| 689 |
-
code = getattr(err, "code", "")
|
| 690 |
-
|
| 691 |
-
if code == "conversation_already_has_active_response":
|
| 692 |
-
# response.create was rejected. The sender worker
|
| 693 |
-
# is waiting on _response_done_event; when the active
|
| 694 |
-
# response finishes it will wake up and see this flag.
|
| 695 |
-
self._last_response_rejected = True
|
| 696 |
-
logger.debug("response.create rejected; worker will retry after active response finishes")
|
| 697 |
-
else:
|
| 698 |
-
logger.error("Realtime error [%s]: %s (raw=%s)", code, msg, err)
|
| 699 |
-
|
| 700 |
-
# Only show user-facing errors, not internal state errors
|
| 701 |
-
if code not in ("input_audio_buffer_commit_empty",):
|
| 702 |
-
await self.output_queue.put(
|
| 703 |
-
AdditionalOutputs({"role": "assistant", "content": f"[error] {msg}"})
|
| 704 |
-
)
|
| 705 |
-
finally:
|
| 706 |
-
# Stop the response sender worker.
|
| 707 |
-
if response_sender_task is not None:
|
| 708 |
-
response_sender_task.cancel()
|
| 709 |
-
try:
|
| 710 |
-
await response_sender_task
|
| 711 |
-
except asyncio.CancelledError:
|
| 712 |
-
pass
|
| 713 |
-
|
| 714 |
-
# Stop background tool manager tasks (listener + cleanup) in all patus.
|
| 715 |
-
await self.tool_manager.shutdown()
|
| 716 |
-
|
| 717 |
-
# Microphone receive
|
| 718 |
-
async def receive(self, frame: Tuple[int, NDArray[np.int16]]) -> None:
|
| 719 |
-
"""Receive audio frame from the microphone and send it to the OpenAI server.
|
| 720 |
-
|
| 721 |
-
Handles both mono and stereo audio formats, converting to the expected
|
| 722 |
-
mono format for OpenAI's API. Resamples if the input sample rate differs
|
| 723 |
-
from the expected rate.
|
| 724 |
-
|
| 725 |
-
Args:
|
| 726 |
-
frame: A tuple containing (sample_rate, audio_data).
|
| 727 |
-
|
| 728 |
-
"""
|
| 729 |
-
if not self.connection:
|
| 730 |
-
return
|
| 731 |
-
|
| 732 |
-
input_sample_rate, audio_frame = frame
|
| 733 |
-
|
| 734 |
-
# Reshape if needed
|
| 735 |
-
if audio_frame.ndim == 2:
|
| 736 |
-
# Scipy channels last convention
|
| 737 |
-
if audio_frame.shape[1] > audio_frame.shape[0]:
|
| 738 |
-
audio_frame = audio_frame.T
|
| 739 |
-
# Multiple channels -> Mono channel
|
| 740 |
-
if audio_frame.shape[1] > 1:
|
| 741 |
-
audio_frame = audio_frame[:, 0]
|
| 742 |
-
|
| 743 |
-
# Resample if needed
|
| 744 |
-
if self.input_sample_rate != input_sample_rate:
|
| 745 |
-
audio_frame = resample(audio_frame, int(len(audio_frame) * self.input_sample_rate / input_sample_rate))
|
| 746 |
-
|
| 747 |
-
# Cast if needed
|
| 748 |
-
audio_frame = audio_to_int16(audio_frame)
|
| 749 |
-
|
| 750 |
-
# Send to OpenAI (guard against races during reconnect)
|
| 751 |
-
try:
|
| 752 |
-
audio_message = base64.b64encode(audio_frame.tobytes()).decode("utf-8")
|
| 753 |
-
await self.connection.input_audio_buffer.append(audio=audio_message)
|
| 754 |
except Exception as e:
|
| 755 |
-
|
| 756 |
-
|
| 757 |
-
|
| 758 |
-
async def emit(self) -> Tuple[int, NDArray[np.int16]] | AdditionalOutputs | None:
|
| 759 |
-
"""Emit audio frame to be played by the speaker."""
|
| 760 |
-
# sends to the stream the stuff put in the output queue by the openai event handler
|
| 761 |
-
# This is called periodically by the fastrtc Stream
|
| 762 |
-
|
| 763 |
-
# Handle idle
|
| 764 |
-
idle_duration = asyncio.get_event_loop().time() - self.last_activity_time
|
| 765 |
-
if idle_duration > 15.0 and self.deps.movement_manager.is_idle():
|
| 766 |
-
try:
|
| 767 |
-
await self.send_idle_signal(idle_duration)
|
| 768 |
-
except Exception as e:
|
| 769 |
-
logger.warning("Idle signal skipped (connection closed?): %s", e)
|
| 770 |
-
return None
|
| 771 |
-
|
| 772 |
-
self.last_activity_time = asyncio.get_event_loop().time() # avoid repeated resets
|
| 773 |
-
|
| 774 |
-
return await wait_for_item(self.output_queue) # type: ignore[no-any-return]
|
| 775 |
-
|
| 776 |
-
async def shutdown(self) -> None:
|
| 777 |
-
"""Shutdown the handler."""
|
| 778 |
-
# Unblock the response sender worker so it can exit
|
| 779 |
-
self._response_done_event.set()
|
| 780 |
-
|
| 781 |
-
# Stop background tool manager tasks (listener + cleanup)
|
| 782 |
-
await self.tool_manager.shutdown()
|
| 783 |
-
|
| 784 |
-
# Cancel any pending debounce task
|
| 785 |
-
if self.partial_transcript_task and not self.partial_transcript_task.done():
|
| 786 |
-
self.partial_transcript_task.cancel()
|
| 787 |
-
try:
|
| 788 |
-
await self.partial_transcript_task
|
| 789 |
-
except asyncio.CancelledError:
|
| 790 |
-
pass
|
| 791 |
-
|
| 792 |
-
if self.connection:
|
| 793 |
-
try:
|
| 794 |
-
await self.connection.close()
|
| 795 |
-
except ConnectionClosedError as e:
|
| 796 |
-
logger.debug(f"Connection already closed during shutdown: {e}")
|
| 797 |
-
except Exception as e:
|
| 798 |
-
logger.debug(f"connection.close() ignored: {e}")
|
| 799 |
-
finally:
|
| 800 |
-
self.connection = None
|
| 801 |
-
|
| 802 |
-
# Clear any remaining items in the output queue
|
| 803 |
-
while not self.output_queue.empty():
|
| 804 |
-
try:
|
| 805 |
-
self.output_queue.get_nowait()
|
| 806 |
-
except asyncio.QueueEmpty:
|
| 807 |
-
break
|
| 808 |
|
| 809 |
-
def
|
| 810 |
-
"""
|
| 811 |
-
|
| 812 |
-
|
| 813 |
-
|
| 814 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 815 |
|
| 816 |
async def get_available_voices(self) -> list[str]:
|
| 817 |
-
"""Try to discover available voices for the configured realtime model.
|
| 818 |
|
| 819 |
Attempts to retrieve model metadata from the OpenAI Models API and look
|
| 820 |
for any keys that might contain voice names. Falls back to a curated
|
| 821 |
list known to work with realtime if discovery fails.
|
| 822 |
"""
|
| 823 |
-
fallback =
|
| 824 |
try:
|
| 825 |
-
# Best effort discovery; safe-guarded for unexpected shapes
|
| 826 |
model = await self.client.models.retrieve(config.MODEL_NAME)
|
| 827 |
-
# Try common serialization paths
|
| 828 |
raw = None
|
| 829 |
for attr in ("model_dump", "to_dict"):
|
| 830 |
fn = getattr(model, attr, None)
|
|
@@ -839,123 +176,46 @@ class OpenaiRealtimeHandler(AsyncStreamHandler):
|
|
| 839 |
raw = dict(model)
|
| 840 |
except Exception:
|
| 841 |
raw = None
|
| 842 |
-
|
| 843 |
candidates: set[str] = set()
|
| 844 |
|
| 845 |
def _collect(obj: object) -> None:
|
| 846 |
try:
|
| 847 |
if isinstance(obj, dict):
|
| 848 |
-
for
|
| 849 |
-
|
| 850 |
-
if "voice" in
|
| 851 |
-
for item in
|
| 852 |
if isinstance(item, str):
|
| 853 |
candidates.add(item)
|
| 854 |
-
elif isinstance(item, dict) and
|
| 855 |
candidates.add(item["name"])
|
| 856 |
else:
|
| 857 |
-
_collect(
|
| 858 |
elif isinstance(obj, (list, tuple)):
|
| 859 |
-
for
|
| 860 |
-
_collect(
|
| 861 |
except Exception:
|
| 862 |
pass
|
| 863 |
|
| 864 |
if isinstance(raw, dict):
|
| 865 |
_collect(raw)
|
| 866 |
-
|
| 867 |
voices = sorted(candidates) if candidates else fallback
|
| 868 |
-
|
| 869 |
-
|
|
|
|
| 870 |
return voices
|
| 871 |
except Exception:
|
| 872 |
return fallback
|
| 873 |
|
| 874 |
-
async def
|
| 875 |
-
"""
|
| 876 |
-
|
| 877 |
-
self.
|
| 878 |
-
|
| 879 |
-
|
| 880 |
-
|
| 881 |
-
|
| 882 |
-
|
| 883 |
-
|
| 884 |
-
"type": "message",
|
| 885 |
-
"role": "user",
|
| 886 |
-
"content": [{"type": "input_text", "text": timestamp_msg}],
|
| 887 |
-
},
|
| 888 |
-
)
|
| 889 |
-
await self._safe_response_create(
|
| 890 |
-
response=RealtimeResponseCreateParamsParam(
|
| 891 |
-
instructions="You MUST respond with function calls only - no speech or text. Choose appropriate actions for idle behavior.",
|
| 892 |
-
tool_choice="required",
|
| 893 |
-
),
|
| 894 |
-
)
|
| 895 |
-
|
| 896 |
-
def _persist_api_key_if_needed(self) -> None:
|
| 897 |
-
"""Persist the API key into `.env` inside `instance_path/` when appropriate.
|
| 898 |
-
|
| 899 |
-
- Only runs in Gradio mode when key came from the textbox and is non-empty.
|
| 900 |
-
- Only saves if `self.instance_path` is not None.
|
| 901 |
-
- Writes `.env` to `instance_path/.env` (does not overwrite if it already exists).
|
| 902 |
-
- If `instance_path/.env.example` exists, copies its contents while overriding OPENAI_API_KEY.
|
| 903 |
-
"""
|
| 904 |
-
try:
|
| 905 |
-
if not self.gradio_mode:
|
| 906 |
-
logger.warning("Not in Gradio mode; skipping API key persistence.")
|
| 907 |
-
return
|
| 908 |
-
|
| 909 |
-
if self._key_source != "textbox":
|
| 910 |
-
logger.info("API key not provided via textbox; skipping persistence.")
|
| 911 |
-
return
|
| 912 |
-
|
| 913 |
-
key = (self._provided_api_key or "").strip()
|
| 914 |
-
if not key:
|
| 915 |
-
logger.warning("No API key provided via textbox; skipping persistence.")
|
| 916 |
-
return
|
| 917 |
-
if self.instance_path is None:
|
| 918 |
-
logger.warning("Instance path is None; cannot persist API key.")
|
| 919 |
-
return
|
| 920 |
-
|
| 921 |
-
# Update the current process environment for downstream consumers
|
| 922 |
-
try:
|
| 923 |
-
import os
|
| 924 |
-
|
| 925 |
-
os.environ["OPENAI_API_KEY"] = key
|
| 926 |
-
except Exception: # best-effort
|
| 927 |
-
pass
|
| 928 |
-
|
| 929 |
-
target_dir = Path(self.instance_path)
|
| 930 |
-
env_path = target_dir / ".env"
|
| 931 |
-
if env_path.exists():
|
| 932 |
-
# Respect existing user configuration
|
| 933 |
-
logger.info(".env already exists at %s; not overwriting.", env_path)
|
| 934 |
-
return
|
| 935 |
-
|
| 936 |
-
example_path = target_dir / ".env.example"
|
| 937 |
-
content_lines: list[str] = []
|
| 938 |
-
if example_path.exists():
|
| 939 |
-
try:
|
| 940 |
-
content = example_path.read_text(encoding="utf-8")
|
| 941 |
-
content_lines = content.splitlines()
|
| 942 |
-
except Exception as e:
|
| 943 |
-
logger.warning("Failed to read .env.example at %s: %s", example_path, e)
|
| 944 |
-
|
| 945 |
-
# Replace or append the OPENAI_API_KEY line
|
| 946 |
-
replaced = False
|
| 947 |
-
for i, line in enumerate(content_lines):
|
| 948 |
-
if line.strip().startswith("OPENAI_API_KEY="):
|
| 949 |
-
content_lines[i] = f"OPENAI_API_KEY={key}"
|
| 950 |
-
replaced = True
|
| 951 |
-
break
|
| 952 |
-
if not replaced:
|
| 953 |
-
content_lines.append(f"OPENAI_API_KEY={key}")
|
| 954 |
-
|
| 955 |
-
# Ensure file ends with newline
|
| 956 |
-
final_text = "\n".join(content_lines) + "\n"
|
| 957 |
-
env_path.write_text(final_text, encoding="utf-8")
|
| 958 |
-
logger.info("Created %s and stored OPENAI_API_KEY for future runs.", env_path)
|
| 959 |
-
except Exception as e:
|
| 960 |
-
# Never crash the app for QoL persistence; just log.
|
| 961 |
-
logger.warning("Could not persist OPENAI_API_KEY to .env: %s", e)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
import logging
|
| 2 |
+
from typing import Any, Literal
|
| 3 |
from pathlib import Path
|
|
|
|
| 4 |
|
|
|
|
|
|
|
| 5 |
from openai import AsyncOpenAI
|
|
|
|
|
|
|
|
|
|
|
|
|
| 6 |
from openai.types.realtime import (
|
| 7 |
AudioTranscriptionParam,
|
| 8 |
RealtimeAudioConfigParam,
|
| 9 |
RealtimeAudioConfigInputParam,
|
| 10 |
RealtimeAudioConfigOutputParam,
|
|
|
|
| 11 |
RealtimeSessionCreateRequestParam,
|
| 12 |
)
|
|
|
|
|
|
|
| 13 |
from openai.types.realtime.realtime_audio_formats_param import AudioPCM
|
| 14 |
from openai.types.realtime.realtime_audio_input_turn_detection_param import ServerVad
|
| 15 |
|
| 16 |
+
from reachy_mini_conversation_app.config import OPENAI_BACKEND, config, get_default_voice_for_backend
|
| 17 |
from reachy_mini_conversation_app.prompts import get_session_voice, get_session_instructions
|
| 18 |
+
from reachy_mini_conversation_app.base_realtime import BaseRealtimeHandler, to_realtime_tools_config
|
| 19 |
+
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies, get_active_tool_specs
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
|
| 21 |
|
| 22 |
logger = logging.getLogger(__name__)
|
| 23 |
|
| 24 |
+
__all__ = ["OpenaiRealtimeHandler"]
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class OpenaiRealtimeHandler(BaseRealtimeHandler):
|
| 28 |
+
"""Realtime handler for the direct OpenAI Realtime API."""
|
| 29 |
+
|
| 30 |
+
BACKEND_PROVIDER = OPENAI_BACKEND
|
| 31 |
+
SAMPLE_RATE = 24000
|
| 32 |
+
REFRESH_CLIENT_ON_RECONNECT = False
|
| 33 |
+
AUDIO_INPUT_COST_PER_1M = 32.0
|
| 34 |
+
AUDIO_OUTPUT_COST_PER_1M = 64.0
|
| 35 |
+
TEXT_INPUT_COST_PER_1M = 4.0
|
| 36 |
+
TEXT_OUTPUT_COST_PER_1M = 16.0
|
| 37 |
+
IMAGE_INPUT_COST_PER_1M = 5.0
|
| 38 |
+
|
| 39 |
+
def __init__(
|
| 40 |
+
self,
|
| 41 |
+
deps: ToolDependencies,
|
| 42 |
+
gradio_mode: bool = False,
|
| 43 |
+
instance_path: str | None = None,
|
| 44 |
+
startup_voice: str | None = None,
|
| 45 |
+
) -> None:
|
| 46 |
+
"""Initialize OpenAI-specific credential state."""
|
| 47 |
+
super().__init__(deps, gradio_mode, instance_path, startup_voice)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 48 |
self._key_source: Literal["env", "textbox"] = "env"
|
| 49 |
self._provided_api_key: str | None = None
|
| 50 |
|
| 51 |
+
async def _prepare_startup_credentials(self) -> None:
|
| 52 |
+
"""Collect an OpenAI API key from Gradio input when needed."""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
openai_api_key = config.OPENAI_API_KEY
|
| 54 |
+
if not self.gradio_mode or openai_api_key:
|
| 55 |
+
return
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 56 |
|
| 57 |
+
await self.wait_for_args() # type: ignore[no-untyped-call]
|
| 58 |
+
args = list(self.latest_args)
|
| 59 |
+
textbox_api_key = args[3] if len(args) > 3 and len(args[3]) > 0 else None
|
| 60 |
+
if textbox_api_key is not None:
|
| 61 |
+
self._key_source = "textbox"
|
| 62 |
+
self._provided_api_key = textbox_api_key
|
| 63 |
|
| 64 |
+
def _persist_credentials_if_needed(self) -> None:
|
| 65 |
+
"""Persist a textbox-provided OpenAI API key into the instance `.env`."""
|
| 66 |
+
try:
|
| 67 |
+
if not self.gradio_mode:
|
| 68 |
+
logger.warning("Not in Gradio mode; skipping OpenAI API key persistence.")
|
| 69 |
return
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 70 |
|
| 71 |
+
if self._key_source != "textbox":
|
| 72 |
+
logger.info("OpenAI API key not provided via textbox; skipping persistence.")
|
| 73 |
+
return
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 74 |
|
| 75 |
+
key = (self._provided_api_key or "").strip()
|
| 76 |
+
if not key:
|
| 77 |
+
logger.warning("No OpenAI API key provided via textbox; skipping persistence.")
|
| 78 |
+
return
|
| 79 |
+
if self.instance_path is None:
|
| 80 |
+
logger.warning("Instance path is None; cannot persist OpenAI API key.")
|
| 81 |
return
|
| 82 |
|
| 83 |
+
# Update the current process environment for downstream consumers.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
try:
|
| 85 |
+
import os
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 86 |
|
| 87 |
+
os.environ["OPENAI_API_KEY"] = key
|
| 88 |
+
except Exception: # best-effort
|
| 89 |
+
pass
|
| 90 |
|
| 91 |
+
target_dir = Path(self.instance_path)
|
| 92 |
+
env_path = target_dir / ".env"
|
| 93 |
+
if env_path.exists():
|
| 94 |
+
# Respect existing user configuration.
|
| 95 |
+
logger.info(".env already exists at %s; not overwriting.", env_path)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
return
|
| 97 |
|
| 98 |
+
example_path = target_dir / ".env.example"
|
| 99 |
+
content_lines: list[str] = []
|
| 100 |
+
if example_path.exists():
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
try:
|
| 102 |
+
content = example_path.read_text(encoding="utf-8")
|
| 103 |
+
content_lines = content.splitlines()
|
| 104 |
except Exception as e:
|
| 105 |
+
logger.warning("Failed to read .env.example at %s: %s", example_path, e)
|
|
|
|
|
|
|
| 106 |
|
| 107 |
+
replaced = False
|
| 108 |
+
for i, line in enumerate(content_lines):
|
| 109 |
+
if line.strip().startswith("OPENAI_API_KEY="):
|
| 110 |
+
content_lines[i] = f"OPENAI_API_KEY={key}"
|
| 111 |
+
replaced = True
|
| 112 |
break
|
| 113 |
+
if not replaced:
|
| 114 |
+
content_lines.append(f"OPENAI_API_KEY={key}")
|
| 115 |
|
| 116 |
+
final_text = "\n".join(content_lines) + "\n"
|
| 117 |
+
env_path.write_text(final_text, encoding="utf-8")
|
| 118 |
+
logger.info("Created %s and stored OPENAI_API_KEY for future runs.", env_path)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 119 |
except Exception as e:
|
| 120 |
+
# Never crash the app for QoL persistence; just log.
|
| 121 |
+
logger.warning("Could not persist OPENAI_API_KEY to .env: %s", e)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 122 |
|
| 123 |
+
def _get_session_instructions(self) -> str:
|
| 124 |
+
"""Return OpenAI session instructions."""
|
| 125 |
+
return get_session_instructions()
|
| 126 |
+
|
| 127 |
+
def _get_session_voice(self, default: str | None = None) -> str:
|
| 128 |
+
"""Return the configured OpenAI session voice."""
|
| 129 |
+
return get_session_voice(default)
|
| 130 |
+
|
| 131 |
+
def _get_active_tool_specs(self) -> list[dict[str, Any]]:
|
| 132 |
+
"""Return active tool specs for the current session dependencies."""
|
| 133 |
+
return get_active_tool_specs(self.deps)
|
| 134 |
+
|
| 135 |
+
def _get_session_config(self, tool_specs: list[dict[str, Any]]) -> RealtimeSessionCreateRequestParam:
|
| 136 |
+
"""Return the OpenAI Realtime session config."""
|
| 137 |
+
return RealtimeSessionCreateRequestParam(
|
| 138 |
+
type="realtime",
|
| 139 |
+
instructions=self._get_session_instructions(),
|
| 140 |
+
audio=RealtimeAudioConfigParam(
|
| 141 |
+
input=RealtimeAudioConfigInputParam(
|
| 142 |
+
format=AudioPCM(type="audio/pcm", rate=24000),
|
| 143 |
+
transcription=AudioTranscriptionParam(model="gpt-4o-transcribe", language="en"),
|
| 144 |
+
turn_detection=ServerVad(type="server_vad", interrupt_response=True),
|
| 145 |
+
),
|
| 146 |
+
output=RealtimeAudioConfigOutputParam(
|
| 147 |
+
format=AudioPCM(type="audio/pcm", rate=24000),
|
| 148 |
+
voice=self.get_current_voice(),
|
| 149 |
+
),
|
| 150 |
+
),
|
| 151 |
+
tools=to_realtime_tools_config(tool_specs),
|
| 152 |
+
tool_choice="auto",
|
| 153 |
+
)
|
| 154 |
|
| 155 |
async def get_available_voices(self) -> list[str]:
|
| 156 |
+
"""Try to discover available voices for the configured OpenAI realtime model.
|
| 157 |
|
| 158 |
Attempts to retrieve model metadata from the OpenAI Models API and look
|
| 159 |
for any keys that might contain voice names. Falls back to a curated
|
| 160 |
list known to work with realtime if discovery fails.
|
| 161 |
"""
|
| 162 |
+
fallback = await super().get_available_voices()
|
| 163 |
try:
|
|
|
|
| 164 |
model = await self.client.models.retrieve(config.MODEL_NAME)
|
|
|
|
| 165 |
raw = None
|
| 166 |
for attr in ("model_dump", "to_dict"):
|
| 167 |
fn = getattr(model, attr, None)
|
|
|
|
| 176 |
raw = dict(model)
|
| 177 |
except Exception:
|
| 178 |
raw = None
|
| 179 |
+
|
| 180 |
candidates: set[str] = set()
|
| 181 |
|
| 182 |
def _collect(obj: object) -> None:
|
| 183 |
try:
|
| 184 |
if isinstance(obj, dict):
|
| 185 |
+
for key, value in obj.items():
|
| 186 |
+
key_lower = str(key).lower()
|
| 187 |
+
if "voice" in key_lower and isinstance(value, (list, tuple)):
|
| 188 |
+
for item in value:
|
| 189 |
if isinstance(item, str):
|
| 190 |
candidates.add(item)
|
| 191 |
+
elif isinstance(item, dict) and isinstance(item.get("name"), str):
|
| 192 |
candidates.add(item["name"])
|
| 193 |
else:
|
| 194 |
+
_collect(value)
|
| 195 |
elif isinstance(obj, (list, tuple)):
|
| 196 |
+
for item in obj:
|
| 197 |
+
_collect(item)
|
| 198 |
except Exception:
|
| 199 |
pass
|
| 200 |
|
| 201 |
if isinstance(raw, dict):
|
| 202 |
_collect(raw)
|
| 203 |
+
|
| 204 |
voices = sorted(candidates) if candidates else fallback
|
| 205 |
+
default_voice = get_default_voice_for_backend(self.BACKEND_PROVIDER)
|
| 206 |
+
if default_voice not in voices:
|
| 207 |
+
voices = [default_voice, *[voice for voice in voices if voice != default_voice]]
|
| 208 |
return voices
|
| 209 |
except Exception:
|
| 210 |
return fallback
|
| 211 |
|
| 212 |
+
async def _build_realtime_client(self) -> AsyncOpenAI:
|
| 213 |
+
"""Build the OpenAI realtime SDK client."""
|
| 214 |
+
self._realtime_connect_query = {}
|
| 215 |
+
resolved_api_key = (self._provided_api_key or config.OPENAI_API_KEY or "").strip()
|
| 216 |
+
if not resolved_api_key:
|
| 217 |
+
# In headless console mode, LocalStream blocks startup until the key is provided.
|
| 218 |
+
# Unit tests may invoke this handler directly with a stubbed client.
|
| 219 |
+
logger.warning("OPENAI_API_KEY missing. Proceeding with a placeholder (tests/offline).")
|
| 220 |
+
resolved_api_key = "DUMMY"
|
| 221 |
+
return AsyncOpenAI(api_key=resolved_api_key)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
src/reachy_mini_conversation_app/startup_settings.py
ADDED
|
@@ -0,0 +1,106 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Helpers for persisting UI-selected startup profile and voice settings."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
import os
|
| 5 |
+
import json
|
| 6 |
+
import logging
|
| 7 |
+
from pathlib import Path
|
| 8 |
+
from dataclasses import dataclass
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
logger = logging.getLogger(__name__)
|
| 12 |
+
|
| 13 |
+
STARTUP_SETTINGS_FILENAME = "startup_settings.json"
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
@dataclass(frozen=True)
|
| 17 |
+
class StartupSettings:
|
| 18 |
+
"""Instance-local startup profile/voice settings selected from the UI."""
|
| 19 |
+
|
| 20 |
+
profile: str | None = None
|
| 21 |
+
voice: str | None = None
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def _normalize_optional_text(value: object) -> str | None:
|
| 25 |
+
"""Return a stripped string or None for empty/non-string values."""
|
| 26 |
+
if not isinstance(value, str):
|
| 27 |
+
return None
|
| 28 |
+
normalized = value.strip()
|
| 29 |
+
return normalized or None
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def _startup_settings_path(instance_path: str | Path | None) -> Path | None:
|
| 33 |
+
"""Return the startup settings JSON path for an instance directory."""
|
| 34 |
+
if instance_path is None:
|
| 35 |
+
return None
|
| 36 |
+
return Path(instance_path) / STARTUP_SETTINGS_FILENAME
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def read_startup_settings(instance_path: str | Path | None) -> StartupSettings:
|
| 40 |
+
"""Read startup settings from an instance-local JSON file."""
|
| 41 |
+
settings_path = _startup_settings_path(instance_path)
|
| 42 |
+
if settings_path is None or not settings_path.exists():
|
| 43 |
+
return StartupSettings()
|
| 44 |
+
|
| 45 |
+
try:
|
| 46 |
+
payload = json.loads(settings_path.read_text(encoding="utf-8"))
|
| 47 |
+
except Exception as exc:
|
| 48 |
+
logger.warning("Failed to read startup settings from %s: %s", settings_path, exc)
|
| 49 |
+
return StartupSettings()
|
| 50 |
+
|
| 51 |
+
if not isinstance(payload, dict):
|
| 52 |
+
logger.warning("Ignoring invalid startup settings payload from %s: %r", settings_path, payload)
|
| 53 |
+
return StartupSettings()
|
| 54 |
+
|
| 55 |
+
return StartupSettings(
|
| 56 |
+
profile=_normalize_optional_text(payload.get("profile")),
|
| 57 |
+
voice=_normalize_optional_text(payload.get("voice")),
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def write_startup_settings(
|
| 62 |
+
instance_path: str | Path | None,
|
| 63 |
+
*,
|
| 64 |
+
profile: str | None,
|
| 65 |
+
voice: str | None,
|
| 66 |
+
) -> None:
|
| 67 |
+
"""Persist startup settings in an instance-local JSON file."""
|
| 68 |
+
settings_path = _startup_settings_path(instance_path)
|
| 69 |
+
if settings_path is None:
|
| 70 |
+
return
|
| 71 |
+
|
| 72 |
+
settings = StartupSettings(
|
| 73 |
+
profile=_normalize_optional_text(profile),
|
| 74 |
+
voice=_normalize_optional_text(voice),
|
| 75 |
+
)
|
| 76 |
+
if settings.profile is None and settings.voice is None:
|
| 77 |
+
try:
|
| 78 |
+
settings_path.unlink()
|
| 79 |
+
except FileNotFoundError:
|
| 80 |
+
return
|
| 81 |
+
return
|
| 82 |
+
|
| 83 |
+
payload: dict[str, str] = {}
|
| 84 |
+
if settings.profile is not None:
|
| 85 |
+
payload["profile"] = settings.profile
|
| 86 |
+
if settings.voice is not None:
|
| 87 |
+
payload["voice"] = settings.voice
|
| 88 |
+
|
| 89 |
+
settings_path.write_text(f"{json.dumps(payload, indent=2, sort_keys=True)}\n", encoding="utf-8")
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
def load_startup_settings_into_runtime(instance_path: str | Path | None) -> StartupSettings:
|
| 93 |
+
"""Load instance-local startup settings when no explicit profile override is set."""
|
| 94 |
+
from reachy_mini_conversation_app.config import LOCKED_PROFILE, set_custom_profile
|
| 95 |
+
|
| 96 |
+
if LOCKED_PROFILE is not None:
|
| 97 |
+
return StartupSettings()
|
| 98 |
+
|
| 99 |
+
settings_path = _startup_settings_path(instance_path)
|
| 100 |
+
settings = read_startup_settings(instance_path)
|
| 101 |
+
if settings_path is None or not settings_path.exists():
|
| 102 |
+
if os.getenv("REACHY_MINI_CUSTOM_PROFILE"):
|
| 103 |
+
return StartupSettings(voice=settings.voice)
|
| 104 |
+
|
| 105 |
+
set_custom_profile(settings.profile)
|
| 106 |
+
return settings
|
src/reachy_mini_conversation_app/static/index.html
CHANGED
|
@@ -16,18 +16,18 @@
|
|
| 16 |
<header class="hero">
|
| 17 |
<div class="pill">Headless control</div>
|
| 18 |
<h1>Reachy Mini Conversation</h1>
|
| 19 |
-
<p class="subtitle">Choose a realtime backend, add
|
| 20 |
</header>
|
| 21 |
|
| 22 |
<aside class="privacy-notice" role="note" aria-label="Privacy notice">
|
| 23 |
<p>
|
| 24 |
<strong>Privacy notice:</strong>
|
| 25 |
-
Speech and camera data are sent to your selected realtime backend for inference. See OpenAI's
|
| 26 |
<a href="https://platform.openai.com/docs/guides/your-data" target="_blank" rel="noopener">
|
| 27 |
data usage policy
|
| 28 |
</a>
|
| 29 |
and Google's
|
| 30 |
-
<a href="https://ai.google.dev/gemini-api/
|
| 31 |
Gemini API data usage policy
|
| 32 |
</a>
|
| 33 |
for details.
|
|
@@ -37,11 +37,13 @@
|
|
| 37 |
<summary>Learn more</summary>
|
| 38 |
|
| 39 |
<p>
|
| 40 |
-
In privacy-sensitive environments, use
|
|
|
|
|
|
|
| 41 |
<a href="https://github.com/pollen-robotics/reachy_mini_conversation_app/blob/main/README.md" target="_blank" rel="noopener">
|
| 42 |
project README
|
| 43 |
</a>
|
| 44 |
-
for full setup instructions.
|
| 45 |
</p>
|
| 46 |
</details>
|
| 47 |
</aside>
|
|
@@ -55,17 +57,24 @@
|
|
| 55 |
<span id="backend-chip" class="chip">Required</span>
|
| 56 |
</div>
|
| 57 |
<p id="backend-note" class="muted backend-note">
|
| 58 |
-
OpenAI Realtime
|
| 59 |
</p>
|
| 60 |
<p class="muted backend-note">
|
| 61 |
Backend changes are saved here, then applied after you restart Reachy Mini Conversation from the dashboard or desktop app.
|
| 62 |
</p>
|
| 63 |
<div class="backend-choice-grid" role="radiogroup" aria-label="Conversation backend">
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 64 |
<label class="backend-choice is-selected" data-backend-card="openai">
|
| 65 |
<input type="radio" name="backend" value="openai" checked />
|
| 66 |
<span class="backend-choice-body">
|
| 67 |
<span class="backend-choice-title">OpenAI Realtime</span>
|
| 68 |
-
<span class="backend-choice-copy">
|
| 69 |
</span>
|
| 70 |
</label>
|
| 71 |
<label class="backend-choice" data-backend-card="gemini">
|
|
@@ -103,8 +112,43 @@
|
|
| 103 |
<span class="chip">Required</span>
|
| 104 |
</div>
|
| 105 |
<p id="form-copy" class="muted">Paste your API key once and we will store it locally for the headless conversation loop.</p>
|
| 106 |
-
<
|
| 107 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 108 |
<div class="actions">
|
| 109 |
<button id="save-btn">Save key</button>
|
| 110 |
<p id="status" class="status" role="status" aria-live="polite" aria-atomic="true"></p>
|
|
|
|
| 16 |
<header class="hero">
|
| 17 |
<div class="pill">Headless control</div>
|
| 18 |
<h1>Reachy Mini Conversation</h1>
|
| 19 |
+
<p class="subtitle">Choose a realtime backend, add credentials only when needed, and tweak personalities without the full UI.</p>
|
| 20 |
</header>
|
| 21 |
|
| 22 |
<aside class="privacy-notice" role="note" aria-label="Privacy notice">
|
| 23 |
<p>
|
| 24 |
<strong>Privacy notice:</strong>
|
| 25 |
+
Speech and camera data are sent to your selected realtime backend for inference. Hugging Face won't store any speech or images sent to our server. See OpenAI's
|
| 26 |
<a href="https://platform.openai.com/docs/guides/your-data" target="_blank" rel="noopener">
|
| 27 |
data usage policy
|
| 28 |
</a>
|
| 29 |
and Google's
|
| 30 |
+
<a href="https://ai.google.dev/gemini-api/terms#data-use-unpaid" target="_blank" rel="noopener">
|
| 31 |
Gemini API data usage policy
|
| 32 |
</a>
|
| 33 |
for details.
|
|
|
|
| 37 |
<summary>Learn more</summary>
|
| 38 |
|
| 39 |
<p>
|
| 40 |
+
In privacy-sensitive environments, use your own local version of
|
| 41 |
+
<a href="https://github.com/huggingface/speech-to-speech" target="_blank" rel="noopener">speech-to-speech</a>,
|
| 42 |
+
or combine any backend with <code>--local-vision</code> to process images on-device. See the
|
| 43 |
<a href="https://github.com/pollen-robotics/reachy_mini_conversation_app/blob/main/README.md" target="_blank" rel="noopener">
|
| 44 |
project README
|
| 45 |
</a>
|
| 46 |
+
for full setup instructions.
|
| 47 |
</p>
|
| 48 |
</details>
|
| 49 |
</aside>
|
|
|
|
| 57 |
<span id="backend-chip" class="chip">Required</span>
|
| 58 |
</div>
|
| 59 |
<p id="backend-note" class="muted backend-note">
|
| 60 |
+
OpenAI Realtime needs your own <code>OPENAI_API_KEY</code>. Gemini Live needs your own <code>GEMINI_API_KEY</code>. Hugging Face uses either the built-in server or a direct realtime websocket endpoint.
|
| 61 |
</p>
|
| 62 |
<p class="muted backend-note">
|
| 63 |
Backend changes are saved here, then applied after you restart Reachy Mini Conversation from the dashboard or desktop app.
|
| 64 |
</p>
|
| 65 |
<div class="backend-choice-grid" role="radiogroup" aria-label="Conversation backend">
|
| 66 |
+
<label class="backend-choice" data-backend-card="huggingface">
|
| 67 |
+
<input type="radio" name="backend" value="huggingface" />
|
| 68 |
+
<span class="backend-choice-body">
|
| 69 |
+
<span class="backend-choice-title">Hugging Face</span>
|
| 70 |
+
<span class="backend-choice-copy">Use the built-in Hugging Face server or your own local endpoint.</span>
|
| 71 |
+
</span>
|
| 72 |
+
</label>
|
| 73 |
<label class="backend-choice is-selected" data-backend-card="openai">
|
| 74 |
<input type="radio" name="backend" value="openai" checked />
|
| 75 |
<span class="backend-choice-body">
|
| 76 |
<span class="backend-choice-title">OpenAI Realtime</span>
|
| 77 |
+
<span class="backend-choice-copy">Bring your own <code>OPENAI_API_KEY</code> to use OpenAI Realtime.</span>
|
| 78 |
</span>
|
| 79 |
</label>
|
| 80 |
<label class="backend-choice" data-backend-card="gemini">
|
|
|
|
| 112 |
<span class="chip">Required</span>
|
| 113 |
</div>
|
| 114 |
<p id="form-copy" class="muted">Paste your API key once and we will store it locally for the headless conversation loop.</p>
|
| 115 |
+
<div id="api-key-fields">
|
| 116 |
+
<label id="api-key-label" for="api-key">OpenAI API Key</label>
|
| 117 |
+
<input id="api-key" type="password" placeholder="sk-..." autocomplete="off" />
|
| 118 |
+
</div>
|
| 119 |
+
<div id="hf-fields" class="hidden">
|
| 120 |
+
<label for="hf-mode">Connection mode</label>
|
| 121 |
+
<select id="hf-mode">
|
| 122 |
+
<option value="local">Local</option>
|
| 123 |
+
<option value="deployed">Hugging Face server</option>
|
| 124 |
+
</select>
|
| 125 |
+
|
| 126 |
+
<div id="hf-direct-fields" class="hf-config-grid">
|
| 127 |
+
<p class="muted hf-help hf-local-note">
|
| 128 |
+
Use your own Hugging Face backend powered by
|
| 129 |
+
<a href="https://github.com/huggingface/speech-to-speech" target="_blank" rel="noopener">speech-to-speech</a>.
|
| 130 |
+
Run it either on the same machine, locally on your network, or hosted on Hugging Face.
|
| 131 |
+
</p>
|
| 132 |
+
<div class="hf-field">
|
| 133 |
+
<label for="hf-host-preset">Host</label>
|
| 134 |
+
<select id="hf-host-preset">
|
| 135 |
+
<option value="localhost">localhost</option>
|
| 136 |
+
<option value="custom">Custom IP or hostname</option>
|
| 137 |
+
</select>
|
| 138 |
+
</div>
|
| 139 |
+
<div class="hf-field">
|
| 140 |
+
<label for="hf-port">Port</label>
|
| 141 |
+
<input id="hf-port" type="text" inputmode="numeric" value="8765" placeholder="8765" autocomplete="off" />
|
| 142 |
+
</div>
|
| 143 |
+
</div>
|
| 144 |
+
|
| 145 |
+
<div id="hf-host-custom-wrap" class="hidden">
|
| 146 |
+
<label for="hf-host-custom">Custom host or IP</label>
|
| 147 |
+
<input id="hf-host-custom" type="text" placeholder="192.168.1.20" autocomplete="off" />
|
| 148 |
+
</div>
|
| 149 |
+
|
| 150 |
+
<p id="hf-preview" class="status"></p>
|
| 151 |
+
</div>
|
| 152 |
<div class="actions">
|
| 153 |
<button id="save-btn">Save key</button>
|
| 154 |
<p id="status" class="status" role="status" aria-live="polite" aria-atomic="true"></p>
|
src/reachy_mini_conversation_app/static/main.js
CHANGED
|
@@ -1,5 +1,9 @@
|
|
| 1 |
const OPENAI_BACKEND = "openai";
|
| 2 |
const GEMINI_BACKEND = "gemini";
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
const BACKEND_META = {
|
| 4 |
[OPENAI_BACKEND]: {
|
| 5 |
label: "OpenAI Realtime",
|
|
@@ -9,10 +13,10 @@ const BACKEND_META = {
|
|
| 9 |
saveButton: "Save key",
|
| 10 |
changeButton: "Change OpenAI key",
|
| 11 |
readyTitle: "OpenAI Realtime ready",
|
| 12 |
-
readyCopy: "OpenAI Realtime is configured.
|
| 13 |
-
formCopy: "
|
| 14 |
-
requiredCredentialsCopy: "OpenAI Realtime
|
| 15 |
-
note: "OpenAI Realtime
|
| 16 |
},
|
| 17 |
[GEMINI_BACKEND]: {
|
| 18 |
label: "Gemini Live",
|
|
@@ -25,12 +29,27 @@ const BACKEND_META = {
|
|
| 25 |
readyCopy: "Gemini Live is configured. Your saved Gemini token is ready to use.",
|
| 26 |
formCopy: "Paste your GEMINI_API_KEY once and we will store it locally for the headless conversation loop.",
|
| 27 |
requiredCredentialsCopy: "Gemini Live requires your own GEMINI_API_KEY before you can switch.",
|
| 28 |
-
note: "OpenAI Realtime
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 29 |
},
|
| 30 |
};
|
| 31 |
|
| 32 |
function backendHasCredentials(status, backend) {
|
| 33 |
-
|
|
|
|
|
|
|
| 34 |
}
|
| 35 |
|
| 36 |
function backendCanProceed(status, backend) {
|
|
@@ -39,17 +58,24 @@ function backendCanProceed(status, backend) {
|
|
| 39 |
? !!status.can_proceed_with_gemini
|
| 40 |
: backendHasCredentials(status, backend);
|
| 41 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 42 |
return status.can_proceed_with_openai !== undefined
|
| 43 |
? !!status.can_proceed_with_openai
|
| 44 |
: backendHasCredentials(status, backend);
|
| 45 |
}
|
| 46 |
|
| 47 |
function backendMeta(backend) {
|
| 48 |
-
return BACKEND_META[backend] || BACKEND_META[
|
| 49 |
}
|
| 50 |
|
| 51 |
function formatBackendNote(text) {
|
| 52 |
-
return text
|
|
|
|
|
|
|
| 53 |
}
|
| 54 |
|
| 55 |
const sleep = (ms) => new Promise((resolve) => setTimeout(resolve, ms));
|
|
@@ -113,8 +139,13 @@ async function validateKey(key) {
|
|
| 113 |
return data;
|
| 114 |
}
|
| 115 |
|
| 116 |
-
async function saveBackendConfig(backend, key = "") {
|
| 117 |
const body = { backend, api_key: key };
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 118 |
const resp = await fetch("/backend_config", {
|
| 119 |
method: "POST",
|
| 120 |
headers: { "Content-Type": "application/json" },
|
|
@@ -246,6 +277,22 @@ function setStatusMessage(el, text, tone = "") {
|
|
| 246 |
el.setAttribute("aria-atomic", "true");
|
| 247 |
}
|
| 248 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 249 |
async function init() {
|
| 250 |
const loading = document.getElementById("loading");
|
| 251 |
show(loading, true);
|
|
@@ -263,10 +310,19 @@ async function init() {
|
|
| 263 |
const personalityPanel = document.getElementById("personality-panel");
|
| 264 |
const formTitle = document.getElementById("form-title");
|
| 265 |
const formCopy = document.getElementById("form-copy");
|
|
|
|
| 266 |
const apiKeyLabel = document.getElementById("api-key-label");
|
| 267 |
const saveBtn = document.getElementById("save-btn");
|
| 268 |
const changeKeyBtn = document.getElementById("change-key-btn");
|
| 269 |
const input = document.getElementById("api-key");
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 270 |
|
| 271 |
// Personality elements
|
| 272 |
const pSelect = document.getElementById("personality-select");
|
|
@@ -287,11 +343,51 @@ async function init() {
|
|
| 287 |
dance: ["stop_dance"],
|
| 288 |
play_emotion: ["stop_emotion"],
|
| 289 |
};
|
| 290 |
-
let selectedBackend =
|
| 291 |
let editingCredentials = false;
|
| 292 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 293 |
function setSelectedBackend(backend) {
|
| 294 |
-
selectedBackend =
|
|
|
|
|
|
|
| 295 |
backendInputs.forEach((radio) => {
|
| 296 |
radio.checked = radio.value === selectedBackend;
|
| 297 |
});
|
|
@@ -301,31 +397,42 @@ async function init() {
|
|
| 301 |
}
|
| 302 |
|
| 303 |
function renderCredentialPanels(status) {
|
| 304 |
-
const persistedBackend = status.backend_provider ||
|
| 305 |
const activeBackend = status.active_backend || persistedBackend;
|
| 306 |
const requiresRestart = !!status.requires_restart;
|
| 307 |
const meta = backendMeta(selectedBackend);
|
| 308 |
const canProceedWithSelectedBackend = backendCanProceed(status, selectedBackend);
|
| 309 |
const selectedMatchesPersisted = selectedBackend === persistedBackend;
|
| 310 |
const selectedMatchesActive = selectedBackend === activeBackend;
|
|
|
|
|
|
|
|
|
|
| 311 |
|
| 312 |
backendChip.textContent = selectedBackend === persistedBackend ? "Saved" : "Selected";
|
| 313 |
backendNote.innerHTML = formatBackendNote(meta.note);
|
| 314 |
|
| 315 |
configuredTitle.textContent = meta.readyTitle;
|
| 316 |
-
configuredCopy.textContent = meta.readyCopy;
|
| 317 |
formTitle.textContent = meta.formTitle;
|
| 318 |
-
formCopy.textContent =
|
|
|
|
|
|
|
|
|
|
|
|
|
| 319 |
apiKeyLabel.textContent = meta.inputLabel;
|
| 320 |
input.placeholder = meta.placeholder;
|
| 321 |
saveBtn.textContent = meta.saveButton;
|
| 322 |
changeKeyBtn.textContent = meta.changeButton;
|
| 323 |
|
| 324 |
show(configuredPanel, canProceedWithSelectedBackend && !editingCredentials);
|
| 325 |
-
show(formPanel, editingCredentials || !canProceedWithSelectedBackend);
|
|
|
|
|
|
|
|
|
|
|
|
|
| 326 |
show(
|
| 327 |
backendSaveBtn,
|
| 328 |
-
canProceedWithSelectedBackend && !selectedMatchesPersisted,
|
| 329 |
);
|
| 330 |
backendSaveBtn.textContent = `Use ${meta.label}`;
|
| 331 |
|
|
@@ -356,17 +463,25 @@ async function init() {
|
|
| 356 |
show(personalityPanel, false);
|
| 357 |
|
| 358 |
const st = (await waitForStatus()) || {
|
| 359 |
-
active_backend:
|
| 360 |
-
backend_provider:
|
| 361 |
has_key: false,
|
| 362 |
has_openai_key: false,
|
| 363 |
has_gemini_key: false,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 364 |
can_proceed: false,
|
| 365 |
can_proceed_with_openai: false,
|
| 366 |
can_proceed_with_gemini: false,
|
|
|
|
| 367 |
requires_restart: false,
|
| 368 |
};
|
| 369 |
-
|
|
|
|
| 370 |
statusEl.textContent = "";
|
| 371 |
renderCredentialPanels(st);
|
| 372 |
|
|
@@ -382,6 +497,23 @@ async function init() {
|
|
| 382 |
input.addEventListener("input", () => {
|
| 383 |
input.classList.remove("error");
|
| 384 |
});
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 385 |
|
| 386 |
backendInputs.forEach((radio) => {
|
| 387 |
radio.addEventListener("change", () => {
|
|
@@ -404,6 +536,59 @@ async function init() {
|
|
| 404 |
});
|
| 405 |
|
| 406 |
saveBtn.addEventListener("click", async () => {
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 407 |
const key = input.value.trim();
|
| 408 |
if (!key) {
|
| 409 |
setStatusMessage(statusEl, "Please enter a valid key.", "warn");
|
|
@@ -424,7 +609,7 @@ async function init() {
|
|
| 424 |
} else {
|
| 425 |
setStatusMessage(statusEl, "Saving Gemini token...", "ok");
|
| 426 |
}
|
| 427 |
-
await saveBackendConfig(selectedBackend, key);
|
| 428 |
setStatusMessage(statusEl, "Saved. Reloading…", "ok");
|
| 429 |
window.location.reload();
|
| 430 |
} catch (e) {
|
|
@@ -443,7 +628,7 @@ async function init() {
|
|
| 443 |
}
|
| 444 |
});
|
| 445 |
|
| 446 |
-
if (!(st.can_proceed ?? backendCanProceed(st, st.backend_provider ||
|
| 447 |
show(loading, false);
|
| 448 |
return;
|
| 449 |
}
|
|
|
|
| 1 |
const OPENAI_BACKEND = "openai";
|
| 2 |
const GEMINI_BACKEND = "gemini";
|
| 3 |
+
const HF_BACKEND = "huggingface";
|
| 4 |
+
const DEFAULT_BACKEND = HF_BACKEND;
|
| 5 |
+
const HF_DEFAULT_HOST = "localhost";
|
| 6 |
+
const HF_DEFAULT_PORT = 8765;
|
| 7 |
const BACKEND_META = {
|
| 8 |
[OPENAI_BACKEND]: {
|
| 9 |
label: "OpenAI Realtime",
|
|
|
|
| 13 |
saveButton: "Save key",
|
| 14 |
changeButton: "Change OpenAI key",
|
| 15 |
readyTitle: "OpenAI Realtime ready",
|
| 16 |
+
readyCopy: "OpenAI Realtime is configured. Your saved OpenAI key is ready to use.",
|
| 17 |
+
formCopy: "Paste your OPENAI_API_KEY once and we will store it locally for the headless conversation loop.",
|
| 18 |
+
requiredCredentialsCopy: "OpenAI Realtime requires your own OPENAI_API_KEY before you can switch.",
|
| 19 |
+
note: "OpenAI Realtime requires your own OPENAI_API_KEY.",
|
| 20 |
},
|
| 21 |
[GEMINI_BACKEND]: {
|
| 22 |
label: "Gemini Live",
|
|
|
|
| 29 |
readyCopy: "Gemini Live is configured. Your saved Gemini token is ready to use.",
|
| 30 |
formCopy: "Paste your GEMINI_API_KEY once and we will store it locally for the headless conversation loop.",
|
| 31 |
requiredCredentialsCopy: "Gemini Live requires your own GEMINI_API_KEY before you can switch.",
|
| 32 |
+
note: "OpenAI Realtime requires OPENAI_API_KEY. Gemini Live needs GEMINI_API_KEY.",
|
| 33 |
+
},
|
| 34 |
+
[HF_BACKEND]: {
|
| 35 |
+
label: "Hugging Face",
|
| 36 |
+
formTitle: "Configure Hugging Face",
|
| 37 |
+
inputLabel: "",
|
| 38 |
+
placeholder: "",
|
| 39 |
+
saveButton: "Save connection",
|
| 40 |
+
changeButton: "Edit connection",
|
| 41 |
+
readyTitle: "Hugging Face ready",
|
| 42 |
+
readyCopy: "Hugging Face is configured. You can jump straight to personalities.",
|
| 43 |
+
formCopy: "Choose where Reachy should connect for Hugging Face.",
|
| 44 |
+
requiredCredentialsCopy: "Set up the Hugging Face connection details before switching.",
|
| 45 |
+
note: "Hugging Face can use the built-in server or your own local realtime websocket.",
|
| 46 |
},
|
| 47 |
};
|
| 48 |
|
| 49 |
function backendHasCredentials(status, backend) {
|
| 50 |
+
if (backend === GEMINI_BACKEND) return !!status.has_gemini_key;
|
| 51 |
+
if (backend === HF_BACKEND) return !!(status.has_hf_connection ?? (status.has_hf_session_url || status.has_hf_ws_url));
|
| 52 |
+
return !!status.has_openai_key;
|
| 53 |
}
|
| 54 |
|
| 55 |
function backendCanProceed(status, backend) {
|
|
|
|
| 58 |
? !!status.can_proceed_with_gemini
|
| 59 |
: backendHasCredentials(status, backend);
|
| 60 |
}
|
| 61 |
+
if (backend === HF_BACKEND) {
|
| 62 |
+
return status.can_proceed_with_hf !== undefined
|
| 63 |
+
? !!status.can_proceed_with_hf
|
| 64 |
+
: backendHasCredentials(status, backend);
|
| 65 |
+
}
|
| 66 |
return status.can_proceed_with_openai !== undefined
|
| 67 |
? !!status.can_proceed_with_openai
|
| 68 |
: backendHasCredentials(status, backend);
|
| 69 |
}
|
| 70 |
|
| 71 |
function backendMeta(backend) {
|
| 72 |
+
return BACKEND_META[backend] || BACKEND_META[DEFAULT_BACKEND];
|
| 73 |
}
|
| 74 |
|
| 75 |
function formatBackendNote(text) {
|
| 76 |
+
return text
|
| 77 |
+
.replace("GEMINI_API_KEY", "<code>GEMINI_API_KEY</code>")
|
| 78 |
+
.replace("HF_REALTIME_WS_URL", "<code>HF_REALTIME_WS_URL</code>");
|
| 79 |
}
|
| 80 |
|
| 81 |
const sleep = (ms) => new Promise((resolve) => setTimeout(resolve, ms));
|
|
|
|
| 139 |
return data;
|
| 140 |
}
|
| 141 |
|
| 142 |
+
async function saveBackendConfig(backend, { key = "", hfMode = "", hfHost = "", hfPort = null } = {}) {
|
| 143 |
const body = { backend, api_key: key };
|
| 144 |
+
if (backend === HF_BACKEND) {
|
| 145 |
+
if (hfMode) body.hf_mode = hfMode;
|
| 146 |
+
if (hfHost) body.hf_host = hfHost;
|
| 147 |
+
if (hfPort !== null && hfPort !== undefined) body.hf_port = hfPort;
|
| 148 |
+
}
|
| 149 |
const resp = await fetch("/backend_config", {
|
| 150 |
method: "POST",
|
| 151 |
headers: { "Content-Type": "application/json" },
|
|
|
|
| 277 |
el.setAttribute("aria-atomic", "true");
|
| 278 |
}
|
| 279 |
|
| 280 |
+
function describeHFConfiguration(status) {
|
| 281 |
+
if (status.hf_connection_mode === "local") {
|
| 282 |
+
const host = status.hf_direct_host || HF_DEFAULT_HOST;
|
| 283 |
+
const port = status.hf_direct_port || HF_DEFAULT_PORT;
|
| 284 |
+
return `Hugging Face will connect directly to ${host}:${port}.`;
|
| 285 |
+
}
|
| 286 |
+
if (status.has_hf_session_url) {
|
| 287 |
+
return "Hugging Face will use the built-in server.";
|
| 288 |
+
}
|
| 289 |
+
return "Choose the Hugging Face server or a local realtime endpoint.";
|
| 290 |
+
}
|
| 291 |
+
|
| 292 |
+
function isLocalHFHost(host) {
|
| 293 |
+
return !host || host === "localhost" || host === "127.0.0.1";
|
| 294 |
+
}
|
| 295 |
+
|
| 296 |
async function init() {
|
| 297 |
const loading = document.getElementById("loading");
|
| 298 |
show(loading, true);
|
|
|
|
| 310 |
const personalityPanel = document.getElementById("personality-panel");
|
| 311 |
const formTitle = document.getElementById("form-title");
|
| 312 |
const formCopy = document.getElementById("form-copy");
|
| 313 |
+
const apiKeyFields = document.getElementById("api-key-fields");
|
| 314 |
const apiKeyLabel = document.getElementById("api-key-label");
|
| 315 |
const saveBtn = document.getElementById("save-btn");
|
| 316 |
const changeKeyBtn = document.getElementById("change-key-btn");
|
| 317 |
const input = document.getElementById("api-key");
|
| 318 |
+
const hfFields = document.getElementById("hf-fields");
|
| 319 |
+
const hfMode = document.getElementById("hf-mode");
|
| 320 |
+
const hfDirectFields = document.getElementById("hf-direct-fields");
|
| 321 |
+
const hfHostPreset = document.getElementById("hf-host-preset");
|
| 322 |
+
const hfHostCustomWrap = document.getElementById("hf-host-custom-wrap");
|
| 323 |
+
const hfHostCustom = document.getElementById("hf-host-custom");
|
| 324 |
+
const hfPort = document.getElementById("hf-port");
|
| 325 |
+
const hfPreview = document.getElementById("hf-preview");
|
| 326 |
|
| 327 |
// Personality elements
|
| 328 |
const pSelect = document.getElementById("personality-select");
|
|
|
|
| 343 |
dance: ["stop_dance"],
|
| 344 |
play_emotion: ["stop_emotion"],
|
| 345 |
};
|
| 346 |
+
let selectedBackend = DEFAULT_BACKEND;
|
| 347 |
let editingCredentials = false;
|
| 348 |
|
| 349 |
+
function resolveHFHost() {
|
| 350 |
+
return hfHostPreset.value === "custom" ? hfHostCustom.value.trim() : HF_DEFAULT_HOST;
|
| 351 |
+
}
|
| 352 |
+
|
| 353 |
+
function updateHFControls() {
|
| 354 |
+
const localMode = hfMode.value !== "deployed";
|
| 355 |
+
const customHost = hfHostPreset.value === "custom";
|
| 356 |
+
show(hfDirectFields, localMode);
|
| 357 |
+
show(hfHostCustomWrap, localMode && customHost);
|
| 358 |
+
|
| 359 |
+
if (!localMode) {
|
| 360 |
+
setStatusMessage(hfPreview, "Hugging Face will use the built-in server.");
|
| 361 |
+
return;
|
| 362 |
+
}
|
| 363 |
+
|
| 364 |
+
const host = resolveHFHost() || "<host>";
|
| 365 |
+
const port = (hfPort.value || String(HF_DEFAULT_PORT)).trim();
|
| 366 |
+
setStatusMessage(hfPreview, `Will save ws://${host}:${port}/v1/realtime`);
|
| 367 |
+
}
|
| 368 |
+
|
| 369 |
+
function populateHFFields(status) {
|
| 370 |
+
const mode = status.hf_connection_mode
|
| 371 |
+
|| (status.has_hf_session_url ? "deployed" : "local");
|
| 372 |
+
const existingHost = status.hf_direct_host || HF_DEFAULT_HOST;
|
| 373 |
+
const existingPort = status.hf_direct_port || HF_DEFAULT_PORT;
|
| 374 |
+
|
| 375 |
+
hfMode.value = mode;
|
| 376 |
+
if (isLocalHFHost(existingHost)) {
|
| 377 |
+
hfHostPreset.value = "localhost";
|
| 378 |
+
hfHostCustom.value = "";
|
| 379 |
+
} else {
|
| 380 |
+
hfHostPreset.value = "custom";
|
| 381 |
+
hfHostCustom.value = existingHost;
|
| 382 |
+
}
|
| 383 |
+
hfPort.value = String(existingPort);
|
| 384 |
+
updateHFControls();
|
| 385 |
+
}
|
| 386 |
+
|
| 387 |
function setSelectedBackend(backend) {
|
| 388 |
+
selectedBackend = [OPENAI_BACKEND, GEMINI_BACKEND, HF_BACKEND].includes(backend)
|
| 389 |
+
? backend
|
| 390 |
+
: DEFAULT_BACKEND;
|
| 391 |
backendInputs.forEach((radio) => {
|
| 392 |
radio.checked = radio.value === selectedBackend;
|
| 393 |
});
|
|
|
|
| 397 |
}
|
| 398 |
|
| 399 |
function renderCredentialPanels(status) {
|
| 400 |
+
const persistedBackend = status.backend_provider || DEFAULT_BACKEND;
|
| 401 |
const activeBackend = status.active_backend || persistedBackend;
|
| 402 |
const requiresRestart = !!status.requires_restart;
|
| 403 |
const meta = backendMeta(selectedBackend);
|
| 404 |
const canProceedWithSelectedBackend = backendCanProceed(status, selectedBackend);
|
| 405 |
const selectedMatchesPersisted = selectedBackend === persistedBackend;
|
| 406 |
const selectedMatchesActive = selectedBackend === activeBackend;
|
| 407 |
+
const usesApiKeyForm = selectedBackend === OPENAI_BACKEND || selectedBackend === GEMINI_BACKEND;
|
| 408 |
+
const usesHFForm = selectedBackend === HF_BACKEND;
|
| 409 |
+
const supportsForm = usesApiKeyForm || usesHFForm;
|
| 410 |
|
| 411 |
backendChip.textContent = selectedBackend === persistedBackend ? "Saved" : "Selected";
|
| 412 |
backendNote.innerHTML = formatBackendNote(meta.note);
|
| 413 |
|
| 414 |
configuredTitle.textContent = meta.readyTitle;
|
| 415 |
+
configuredCopy.textContent = usesHFForm ? describeHFConfiguration(status) : meta.readyCopy;
|
| 416 |
formTitle.textContent = meta.formTitle;
|
| 417 |
+
formCopy.textContent = usesHFForm
|
| 418 |
+
? meta.formCopy
|
| 419 |
+
: canProceedWithSelectedBackend
|
| 420 |
+
? meta.formCopy
|
| 421 |
+
: meta.requiredCredentialsCopy;
|
| 422 |
apiKeyLabel.textContent = meta.inputLabel;
|
| 423 |
input.placeholder = meta.placeholder;
|
| 424 |
saveBtn.textContent = meta.saveButton;
|
| 425 |
changeKeyBtn.textContent = meta.changeButton;
|
| 426 |
|
| 427 |
show(configuredPanel, canProceedWithSelectedBackend && !editingCredentials);
|
| 428 |
+
show(formPanel, supportsForm && (editingCredentials || !canProceedWithSelectedBackend));
|
| 429 |
+
show(apiKeyFields, usesApiKeyForm);
|
| 430 |
+
show(hfFields, usesHFForm);
|
| 431 |
+
if (usesHFForm) updateHFControls();
|
| 432 |
+
show(changeKeyBtn, supportsForm && canProceedWithSelectedBackend && !editingCredentials);
|
| 433 |
show(
|
| 434 |
backendSaveBtn,
|
| 435 |
+
canProceedWithSelectedBackend && !selectedMatchesPersisted && !editingCredentials,
|
| 436 |
);
|
| 437 |
backendSaveBtn.textContent = `Use ${meta.label}`;
|
| 438 |
|
|
|
|
| 463 |
show(personalityPanel, false);
|
| 464 |
|
| 465 |
const st = (await waitForStatus()) || {
|
| 466 |
+
active_backend: DEFAULT_BACKEND,
|
| 467 |
+
backend_provider: DEFAULT_BACKEND,
|
| 468 |
has_key: false,
|
| 469 |
has_openai_key: false,
|
| 470 |
has_gemini_key: false,
|
| 471 |
+
has_hf_session_url: false,
|
| 472 |
+
has_hf_ws_url: false,
|
| 473 |
+
has_hf_connection: false,
|
| 474 |
+
hf_connection_mode: "local",
|
| 475 |
+
hf_direct_host: HF_DEFAULT_HOST,
|
| 476 |
+
hf_direct_port: HF_DEFAULT_PORT,
|
| 477 |
can_proceed: false,
|
| 478 |
can_proceed_with_openai: false,
|
| 479 |
can_proceed_with_gemini: false,
|
| 480 |
+
can_proceed_with_hf: false,
|
| 481 |
requires_restart: false,
|
| 482 |
};
|
| 483 |
+
populateHFFields(st);
|
| 484 |
+
setSelectedBackend(st.backend_provider || DEFAULT_BACKEND);
|
| 485 |
statusEl.textContent = "";
|
| 486 |
renderCredentialPanels(st);
|
| 487 |
|
|
|
|
| 497 |
input.addEventListener("input", () => {
|
| 498 |
input.classList.remove("error");
|
| 499 |
});
|
| 500 |
+
hfHostCustom.addEventListener("input", () => {
|
| 501 |
+
hfHostCustom.classList.remove("error");
|
| 502 |
+
updateHFControls();
|
| 503 |
+
});
|
| 504 |
+
hfPort.addEventListener("input", () => {
|
| 505 |
+
hfPort.classList.remove("error");
|
| 506 |
+
updateHFControls();
|
| 507 |
+
});
|
| 508 |
+
hfMode.addEventListener("change", () => {
|
| 509 |
+
hfHostCustom.classList.remove("error");
|
| 510 |
+
hfPort.classList.remove("error");
|
| 511 |
+
updateHFControls();
|
| 512 |
+
});
|
| 513 |
+
hfHostPreset.addEventListener("change", () => {
|
| 514 |
+
hfHostCustom.classList.remove("error");
|
| 515 |
+
updateHFControls();
|
| 516 |
+
});
|
| 517 |
|
| 518 |
backendInputs.forEach((radio) => {
|
| 519 |
radio.addEventListener("change", () => {
|
|
|
|
| 536 |
});
|
| 537 |
|
| 538 |
saveBtn.addEventListener("click", async () => {
|
| 539 |
+
if (selectedBackend === HF_BACKEND) {
|
| 540 |
+
const localMode = hfMode.value !== "deployed";
|
| 541 |
+
setStatusMessage(statusEl, "Saving connection...");
|
| 542 |
+
hfHostCustom.classList.remove("error");
|
| 543 |
+
hfPort.classList.remove("error");
|
| 544 |
+
|
| 545 |
+
try {
|
| 546 |
+
if (localMode) {
|
| 547 |
+
const host = resolveHFHost();
|
| 548 |
+
const port = Number.parseInt((hfPort.value || "").trim(), 10);
|
| 549 |
+
if (!host) {
|
| 550 |
+
hfHostCustom.classList.add("error");
|
| 551 |
+
setStatusMessage(statusEl, "Enter a valid host or IP address.", "warn");
|
| 552 |
+
return;
|
| 553 |
+
}
|
| 554 |
+
if (!Number.isInteger(port) || port < 1 || port > 65535) {
|
| 555 |
+
hfPort.classList.add("error");
|
| 556 |
+
setStatusMessage(statusEl, "Enter a valid port between 1 and 65535.", "warn");
|
| 557 |
+
return;
|
| 558 |
+
}
|
| 559 |
+
|
| 560 |
+
await saveBackendConfig(selectedBackend, {
|
| 561 |
+
hfMode: "local",
|
| 562 |
+
hfHost: host,
|
| 563 |
+
hfPort: port,
|
| 564 |
+
});
|
| 565 |
+
} else {
|
| 566 |
+
await saveBackendConfig(selectedBackend, {
|
| 567 |
+
hfMode: "deployed",
|
| 568 |
+
});
|
| 569 |
+
}
|
| 570 |
+
setStatusMessage(statusEl, "Saved. Reloading…", "ok");
|
| 571 |
+
window.location.reload();
|
| 572 |
+
} catch (e) {
|
| 573 |
+
if (e.message === "missing_hf_session_url") {
|
| 574 |
+
setStatusMessage(
|
| 575 |
+
statusEl,
|
| 576 |
+
"The built-in Hugging Face server URL is unavailable. Restart the app and try again.",
|
| 577 |
+
"error",
|
| 578 |
+
);
|
| 579 |
+
} else if (e.message === "empty_hf_host" || e.message === "invalid_hf_host") {
|
| 580 |
+
hfHostCustom.classList.add("error");
|
| 581 |
+
setStatusMessage(statusEl, "Enter a valid host or IP address.", "error");
|
| 582 |
+
} else if (e.message === "invalid_hf_port") {
|
| 583 |
+
hfPort.classList.add("error");
|
| 584 |
+
setStatusMessage(statusEl, "Enter a valid port between 1 and 65535.", "error");
|
| 585 |
+
} else {
|
| 586 |
+
setStatusMessage(statusEl, "Failed to save the Hugging Face connection.", "error");
|
| 587 |
+
}
|
| 588 |
+
}
|
| 589 |
+
return;
|
| 590 |
+
}
|
| 591 |
+
|
| 592 |
const key = input.value.trim();
|
| 593 |
if (!key) {
|
| 594 |
setStatusMessage(statusEl, "Please enter a valid key.", "warn");
|
|
|
|
| 609 |
} else {
|
| 610 |
setStatusMessage(statusEl, "Saving Gemini token...", "ok");
|
| 611 |
}
|
| 612 |
+
await saveBackendConfig(selectedBackend, { key });
|
| 613 |
setStatusMessage(statusEl, "Saved. Reloading…", "ok");
|
| 614 |
window.location.reload();
|
| 615 |
} catch (e) {
|
|
|
|
| 628 |
}
|
| 629 |
});
|
| 630 |
|
| 631 |
+
if (!(st.can_proceed ?? backendCanProceed(st, st.backend_provider || DEFAULT_BACKEND)) || st.requires_restart) {
|
| 632 |
show(loading, false);
|
| 633 |
return;
|
| 634 |
}
|
src/reachy_mini_conversation_app/static/style.css
CHANGED
|
@@ -274,7 +274,7 @@ body {
|
|
| 274 |
border-color: rgba(76, 224, 179, 0.4);
|
| 275 |
}
|
| 276 |
|
| 277 |
-
.hidden { display: none; }
|
| 278 |
label {
|
| 279 |
display: block;
|
| 280 |
margin: 8px 0 6px;
|
|
@@ -350,6 +350,28 @@ button.ghost:hover { border-color: rgba(94, 240, 193, 0.4); }
|
|
| 350 |
.status.ok { color: var(--ok); }
|
| 351 |
.status.warn { color: var(--warn); }
|
| 352 |
.status.error { color: var(--error); }
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 353 |
|
| 354 |
/* Personality layout */
|
| 355 |
.row {
|
|
@@ -439,6 +461,7 @@ button.ghost:hover { border-color: rgba(94, 240, 193, 0.4); }
|
|
| 439 |
.hero h1 { font-size: 26px; }
|
| 440 |
.privacy-notice { padding: 10px 12px; }
|
| 441 |
.backend-choice-grid { grid-template-columns: 1fr; }
|
|
|
|
| 442 |
.row { grid-template-columns: 1fr; }
|
| 443 |
#personality-panel .row:first-of-type { grid-template-columns: 1fr; }
|
| 444 |
button { width: 100%; justify-content: center; }
|
|
|
|
| 274 |
border-color: rgba(76, 224, 179, 0.4);
|
| 275 |
}
|
| 276 |
|
| 277 |
+
.hidden { display: none !important; }
|
| 278 |
label {
|
| 279 |
display: block;
|
| 280 |
margin: 8px 0 6px;
|
|
|
|
| 350 |
.status.ok { color: var(--ok); }
|
| 351 |
.status.warn { color: var(--warn); }
|
| 352 |
.status.error { color: var(--error); }
|
| 353 |
+
.hf-help {
|
| 354 |
+
margin: 8px 0 0;
|
| 355 |
+
font-size: 13px;
|
| 356 |
+
line-height: 1.5;
|
| 357 |
+
}
|
| 358 |
+
.hf-help a {
|
| 359 |
+
color: var(--accent);
|
| 360 |
+
text-underline-offset: 2px;
|
| 361 |
+
}
|
| 362 |
+
.hf-config-grid {
|
| 363 |
+
display: grid;
|
| 364 |
+
grid-template-columns: repeat(2, minmax(0, 1fr));
|
| 365 |
+
gap: 12px;
|
| 366 |
+
margin-top: 10px;
|
| 367 |
+
}
|
| 368 |
+
.hf-local-note {
|
| 369 |
+
grid-column: 1 / -1;
|
| 370 |
+
margin: 0;
|
| 371 |
+
}
|
| 372 |
+
.hf-field {
|
| 373 |
+
min-width: 0;
|
| 374 |
+
}
|
| 375 |
|
| 376 |
/* Personality layout */
|
| 377 |
.row {
|
|
|
|
| 461 |
.hero h1 { font-size: 26px; }
|
| 462 |
.privacy-notice { padding: 10px 12px; }
|
| 463 |
.backend-choice-grid { grid-template-columns: 1fr; }
|
| 464 |
+
.hf-config-grid { grid-template-columns: 1fr; }
|
| 465 |
.row { grid-template-columns: 1fr; }
|
| 466 |
#personality-panel .row:first-of-type { grid-template-columns: 1fr; }
|
| 467 |
button { width: 100%; justify-content: center; }
|
src/reachy_mini_conversation_app/tools/core_tools.py
CHANGED
|
@@ -27,14 +27,6 @@ if TYPE_CHECKING:
|
|
| 27 |
logger = logging.getLogger(__name__)
|
| 28 |
|
| 29 |
|
| 30 |
-
if not logger.handlers:
|
| 31 |
-
handler = logging.StreamHandler()
|
| 32 |
-
formatter = logging.Formatter("%(asctime)s %(levelname)s %(name)s:%(lineno)d | %(message)s")
|
| 33 |
-
handler.setFormatter(formatter)
|
| 34 |
-
logger.addHandler(handler)
|
| 35 |
-
logger.setLevel(logging.INFO)
|
| 36 |
-
|
| 37 |
-
|
| 38 |
ALL_TOOLS: Dict[str, "Tool"] = {}
|
| 39 |
ALL_TOOL_SPECS: List[Dict[str, Any]] = []
|
| 40 |
_TOOLS_INITIALIZED = False
|
|
@@ -290,6 +282,14 @@ def get_tool_specs(exclusion_list: list[str] = []) -> list[Dict[str, Any]]:
|
|
| 290 |
return [spec for spec in ALL_TOOL_SPECS if spec.get("name") not in exclusion_list]
|
| 291 |
|
| 292 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 293 |
# Dispatcher
|
| 294 |
def _safe_load_obj(args_json: str) -> Dict[str, Any]:
|
| 295 |
try:
|
|
|
|
| 27 |
logger = logging.getLogger(__name__)
|
| 28 |
|
| 29 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
ALL_TOOLS: Dict[str, "Tool"] = {}
|
| 31 |
ALL_TOOL_SPECS: List[Dict[str, Any]] = []
|
| 32 |
_TOOLS_INITIALIZED = False
|
|
|
|
| 282 |
return [spec for spec in ALL_TOOL_SPECS if spec.get("name") not in exclusion_list]
|
| 283 |
|
| 284 |
|
| 285 |
+
def get_active_tool_specs(deps: ToolDependencies) -> list[Dict[str, Any]]:
|
| 286 |
+
"""Get tool specs filtered by what the current session deps support."""
|
| 287 |
+
exclusion_list: list[str] = []
|
| 288 |
+
if not (deps.camera_worker and deps.camera_worker.head_tracker):
|
| 289 |
+
exclusion_list.append("head_tracking")
|
| 290 |
+
return get_tool_specs(exclusion_list)
|
| 291 |
+
|
| 292 |
+
|
| 293 |
# Dispatcher
|
| 294 |
def _safe_load_obj(args_json: str) -> Dict[str, Any]:
|
| 295 |
try:
|
src/reachy_mini_conversation_app/tools/dance.py
CHANGED
|
@@ -1,3 +1,4 @@
|
|
|
|
|
| 1 |
import logging
|
| 2 |
from typing import Any, Dict
|
| 3 |
|
|
@@ -18,6 +19,21 @@ except ImportError as e:
|
|
| 18 |
DANCE_AVAILABLE = False
|
| 19 |
|
| 20 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
class Dance(Tool):
|
| 22 |
"""Play a named or random dance move once (or repeat). Non-blocking."""
|
| 23 |
|
|
@@ -28,28 +44,11 @@ class Dance(Tool):
|
|
| 28 |
"properties": {
|
| 29 |
"move": {
|
| 30 |
"type": "string",
|
| 31 |
-
"
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
dizzy_spin: A circular 'dizzy' head motion combining roll and pitch.
|
| 37 |
-
stumble_and_recover: A simulated stumble and recovery with multiple axis movements. Good vibes
|
| 38 |
-
interwoven_spirals: A complex spiral motion using three axes at different frequencies.
|
| 39 |
-
sharp_side_tilt: A sharp, quick side-to-side tilt using a triangle waveform.
|
| 40 |
-
side_peekaboo: A multi-stage peekaboo performance, hiding and peeking to each side.
|
| 41 |
-
yeah_nod: An emphatic two-part yeah nod using transient motions.
|
| 42 |
-
uh_huh_tilt: A combined roll-and-pitch uh-huh gesture of agreement.
|
| 43 |
-
neck_recoil: A quick, transient backward recoil of the neck.
|
| 44 |
-
chin_lead: A forward motion led by the chin, combining translation and pitch.
|
| 45 |
-
groovy_sway_and_roll: A side-to-side sway combined with a corresponding roll for a groovy effect.
|
| 46 |
-
chicken_peck: A sharp, forward, chicken-like pecking motion.
|
| 47 |
-
side_glance_flick: A quick glance to the side that holds, then returns.
|
| 48 |
-
polyrhythm_combo: A 3-beat sway and a 2-beat nod create a polyrhythmic feel.
|
| 49 |
-
grid_snap: A robotic, grid-snapping motion using square waveforms.
|
| 50 |
-
pendulum_swing: A simple, smooth pendulum-like swing using a roll motion.
|
| 51 |
-
jackson_square: Traces a rectangle via a 5-point path, with sharp twitches on arrival at each checkpoint.
|
| 52 |
-
""",
|
| 53 |
},
|
| 54 |
"repeat": {
|
| 55 |
"type": "integer",
|
|
@@ -64,14 +63,15 @@ class Dance(Tool):
|
|
| 64 |
if not DANCE_AVAILABLE:
|
| 65 |
return {"error": "Dance system not available"}
|
| 66 |
|
|
|
|
|
|
|
|
|
|
| 67 |
move_name = kwargs.get("move")
|
| 68 |
repeat = int(kwargs.get("repeat", 1))
|
| 69 |
|
| 70 |
logger.info("Tool call: dance move=%s repeat=%d", move_name, repeat)
|
| 71 |
|
| 72 |
-
if not move_name
|
| 73 |
-
import random
|
| 74 |
-
|
| 75 |
move_name = random.choice(list(AVAILABLE_MOVES.keys()))
|
| 76 |
|
| 77 |
if move_name not in AVAILABLE_MOVES:
|
|
|
|
| 1 |
+
import random
|
| 2 |
import logging
|
| 3 |
from typing import Any, Dict
|
| 4 |
|
|
|
|
| 19 |
DANCE_AVAILABLE = False
|
| 20 |
|
| 21 |
|
| 22 |
+
def get_available_dances_and_descriptions() -> str:
|
| 23 |
+
"""Get formatted list of available dances with descriptions."""
|
| 24 |
+
if not DANCE_AVAILABLE:
|
| 25 |
+
return "Moves not available."
|
| 26 |
+
|
| 27 |
+
if not AVAILABLE_MOVES: # if AVAILABLE_MOVES is empty
|
| 28 |
+
return "Moves not available."
|
| 29 |
+
|
| 30 |
+
output = ""
|
| 31 |
+
for move_name, (func, params, metadata) in AVAILABLE_MOVES.items():
|
| 32 |
+
description = metadata.get("description", "No description available.")
|
| 33 |
+
output += f"{move_name}: {description}\n"
|
| 34 |
+
return output
|
| 35 |
+
|
| 36 |
+
|
| 37 |
class Dance(Tool):
|
| 38 |
"""Play a named or random dance move once (or repeat). Non-blocking."""
|
| 39 |
|
|
|
|
| 44 |
"properties": {
|
| 45 |
"move": {
|
| 46 |
"type": "string",
|
| 47 |
+
"enum": list(AVAILABLE_MOVES.keys() if DANCE_AVAILABLE else []),
|
| 48 |
+
"description": f"""Name of the moves and their descriptions; omit for random.
|
| 49 |
+
Here is a list of the available moves, you MUST only choose from these: \n
|
| 50 |
+
{get_available_dances_and_descriptions()}
|
| 51 |
+
""",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 52 |
},
|
| 53 |
"repeat": {
|
| 54 |
"type": "integer",
|
|
|
|
| 63 |
if not DANCE_AVAILABLE:
|
| 64 |
return {"error": "Dance system not available"}
|
| 65 |
|
| 66 |
+
if not AVAILABLE_MOVES: # if AVAILABLE_MOVES is empty
|
| 67 |
+
return {"error": "No moves currently available"}
|
| 68 |
+
|
| 69 |
move_name = kwargs.get("move")
|
| 70 |
repeat = int(kwargs.get("repeat", 1))
|
| 71 |
|
| 72 |
logger.info("Tool call: dance move=%s repeat=%d", move_name, repeat)
|
| 73 |
|
| 74 |
+
if not move_name:
|
|
|
|
|
|
|
| 75 |
move_name = random.choice(list(AVAILABLE_MOVES.keys()))
|
| 76 |
|
| 77 |
if move_name not in AVAILABLE_MOVES:
|
src/reachy_mini_conversation_app/tools/{do_nothing.py → idle_do_nothing.py}
RENAMED
|
@@ -7,24 +7,27 @@ from reachy_mini_conversation_app.tools.core_tools import Tool, ToolDependencies
|
|
| 7 |
logger = logging.getLogger(__name__)
|
| 8 |
|
| 9 |
|
| 10 |
-
class
|
| 11 |
-
"""
|
| 12 |
|
| 13 |
-
name = "
|
| 14 |
-
description =
|
|
|
|
|
|
|
|
|
|
| 15 |
parameters_schema = {
|
| 16 |
"type": "object",
|
| 17 |
"properties": {
|
| 18 |
"reason": {
|
| 19 |
"type": "string",
|
| 20 |
-
"description": "Optional reason for
|
| 21 |
},
|
| 22 |
},
|
| 23 |
"required": [],
|
| 24 |
}
|
| 25 |
|
| 26 |
async def __call__(self, deps: ToolDependencies, **kwargs: Any) -> Dict[str, Any]:
|
| 27 |
-
"""
|
| 28 |
-
reason = kwargs.get("reason", "
|
| 29 |
-
logger.info("Tool call:
|
| 30 |
-
return {"status": "
|
|
|
|
| 7 |
logger = logging.getLogger(__name__)
|
| 8 |
|
| 9 |
|
| 10 |
+
class IdleDoNothing(Tool):
|
| 11 |
+
"""Explicitly choose no action during an idle turn."""
|
| 12 |
|
| 13 |
+
name = "idle_do_nothing"
|
| 14 |
+
description = (
|
| 15 |
+
"Use only in response to an idle time update when you intentionally want Reachy to stay still and silent "
|
| 16 |
+
"instead of choosing another idle action."
|
| 17 |
+
)
|
| 18 |
parameters_schema = {
|
| 19 |
"type": "object",
|
| 20 |
"properties": {
|
| 21 |
"reason": {
|
| 22 |
"type": "string",
|
| 23 |
+
"description": "Optional reason for staying idle during this idle turn.",
|
| 24 |
},
|
| 25 |
},
|
| 26 |
"required": [],
|
| 27 |
}
|
| 28 |
|
| 29 |
async def __call__(self, deps: ToolDependencies, **kwargs: Any) -> Dict[str, Any]:
|
| 30 |
+
"""Stay still and silent for the current idle turn."""
|
| 31 |
+
reason = kwargs.get("reason", "idle turn")
|
| 32 |
+
logger.info("Tool call: idle_do_nothing reason=%s", reason)
|
| 33 |
+
return {"status": "idle", "reason": reason}
|
src/reachy_mini_conversation_app/tools/play_emotion.py
CHANGED
|
@@ -1,3 +1,4 @@
|
|
|
|
|
| 1 |
import logging
|
| 2 |
from typing import Any, Dict
|
| 3 |
|
|
@@ -27,6 +28,9 @@ def get_available_emotions_and_descriptions() -> str:
|
|
| 27 |
|
| 28 |
try:
|
| 29 |
emotion_names = RECORDED_MOVES.list_moves()
|
|
|
|
|
|
|
|
|
|
| 30 |
output = "Available emotions:\n"
|
| 31 |
for name in emotion_names:
|
| 32 |
description = RECORDED_MOVES.get(name).description
|
|
@@ -46,13 +50,14 @@ class PlayEmotion(Tool):
|
|
| 46 |
"properties": {
|
| 47 |
"emotion": {
|
| 48 |
"type": "string",
|
| 49 |
-
"
|
| 50 |
-
|
|
|
|
| 51 |
{get_available_emotions_and_descriptions()}
|
| 52 |
""",
|
| 53 |
},
|
| 54 |
},
|
| 55 |
-
"required": [
|
| 56 |
}
|
| 57 |
|
| 58 |
async def __call__(self, deps: ToolDependencies, **kwargs: Any) -> Dict[str, Any]:
|
|
@@ -61,14 +66,18 @@ class PlayEmotion(Tool):
|
|
| 61 |
return {"error": "Emotion system not available"}
|
| 62 |
|
| 63 |
emotion_name = kwargs.get("emotion")
|
| 64 |
-
if not emotion_name:
|
| 65 |
-
return {"error": "Emotion name is required"}
|
| 66 |
|
| 67 |
logger.info("Tool call: play_emotion emotion=%s", emotion_name)
|
| 68 |
|
| 69 |
# Check if emotion exists
|
| 70 |
try:
|
| 71 |
emotion_names = RECORDED_MOVES.list_moves()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 72 |
if emotion_name not in emotion_names:
|
| 73 |
return {"error": f"Unknown emotion '{emotion_name}'. Available: {emotion_names}"}
|
| 74 |
|
|
|
|
| 1 |
+
import random
|
| 2 |
import logging
|
| 3 |
from typing import Any, Dict
|
| 4 |
|
|
|
|
| 28 |
|
| 29 |
try:
|
| 30 |
emotion_names = RECORDED_MOVES.list_moves()
|
| 31 |
+
if not emotion_names:
|
| 32 |
+
return "No emotions currently available"
|
| 33 |
+
|
| 34 |
output = "Available emotions:\n"
|
| 35 |
for name in emotion_names:
|
| 36 |
description = RECORDED_MOVES.get(name).description
|
|
|
|
| 50 |
"properties": {
|
| 51 |
"emotion": {
|
| 52 |
"type": "string",
|
| 53 |
+
"enum": list(RECORDED_MOVES.list_moves()) if EMOTION_AVAILABLE else [],
|
| 54 |
+
"description": f"""Name of the emotion to play; omit for random.
|
| 55 |
+
Here is a list of the available emotions, you MUST only choose from these: \n
|
| 56 |
{get_available_emotions_and_descriptions()}
|
| 57 |
""",
|
| 58 |
},
|
| 59 |
},
|
| 60 |
+
"required": [],
|
| 61 |
}
|
| 62 |
|
| 63 |
async def __call__(self, deps: ToolDependencies, **kwargs: Any) -> Dict[str, Any]:
|
|
|
|
| 66 |
return {"error": "Emotion system not available"}
|
| 67 |
|
| 68 |
emotion_name = kwargs.get("emotion")
|
|
|
|
|
|
|
| 69 |
|
| 70 |
logger.info("Tool call: play_emotion emotion=%s", emotion_name)
|
| 71 |
|
| 72 |
# Check if emotion exists
|
| 73 |
try:
|
| 74 |
emotion_names = RECORDED_MOVES.list_moves()
|
| 75 |
+
if not emotion_names:
|
| 76 |
+
return {"error": "No emotions currently available"}
|
| 77 |
+
|
| 78 |
+
if not emotion_name:
|
| 79 |
+
emotion_name = random.choice(emotion_names)
|
| 80 |
+
|
| 81 |
if emotion_name not in emotion_names:
|
| 82 |
return {"error": f"Unknown emotion '{emotion_name}'. Available: {emotion_names}"}
|
| 83 |
|
src/reachy_mini_conversation_app/vision/head_tracking/yolo.py
CHANGED
|
@@ -1,5 +1,7 @@
|
|
| 1 |
from __future__ import annotations
|
| 2 |
import logging
|
|
|
|
|
|
|
| 3 |
|
| 4 |
import numpy as np
|
| 5 |
from numpy.typing import NDArray
|
|
@@ -9,7 +11,7 @@ from reachy_mini_conversation_app.vision.head_tracking import HeadTrackerResult
|
|
| 9 |
|
| 10 |
try:
|
| 11 |
from supervision import Detections
|
| 12 |
-
from ultralytics import YOLO # type: ignore
|
| 13 |
except ImportError as e:
|
| 14 |
raise ImportError(
|
| 15 |
"To use YOLO head tracker, please install the extra dependencies: pip install '.[yolo_vision]'",
|
|
@@ -20,6 +22,14 @@ from huggingface_hub import hf_hub_download
|
|
| 20 |
logger = logging.getLogger(__name__)
|
| 21 |
|
| 22 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 23 |
class YoloHeadTracker:
|
| 24 |
"""Lightweight head tracker using YOLO for face detection."""
|
| 25 |
|
|
@@ -35,7 +45,7 @@ class YoloHeadTracker:
|
|
| 35 |
|
| 36 |
try:
|
| 37 |
model_path = hf_hub_download(repo_id=model_repo, filename=model_filename)
|
| 38 |
-
self.model = YOLO(model_path).to(device)
|
| 39 |
logger.info("YOLO face detection model loaded from %s", model_repo)
|
| 40 |
except Exception as e:
|
| 41 |
logger.error("Failed to load YOLO model: %s", e)
|
|
|
|
| 1 |
from __future__ import annotations
|
| 2 |
import logging
|
| 3 |
+
from typing import Protocol, cast
|
| 4 |
+
from collections.abc import Sequence
|
| 5 |
|
| 6 |
import numpy as np
|
| 7 |
from numpy.typing import NDArray
|
|
|
|
| 11 |
|
| 12 |
try:
|
| 13 |
from supervision import Detections
|
| 14 |
+
from ultralytics import YOLO # type: ignore[attr-defined]
|
| 15 |
except ImportError as e:
|
| 16 |
raise ImportError(
|
| 17 |
"To use YOLO head tracker, please install the extra dependencies: pip install '.[yolo_vision]'",
|
|
|
|
| 22 |
logger = logging.getLogger(__name__)
|
| 23 |
|
| 24 |
|
| 25 |
+
class _YoloModel(Protocol):
|
| 26 |
+
"""Minimal YOLO model interface used by the head tracker."""
|
| 27 |
+
|
| 28 |
+
def __call__(self, source: NDArray[np.uint8], **kwargs: object) -> Sequence[object]: ...
|
| 29 |
+
|
| 30 |
+
def to(self, device: str) -> _YoloModel: ...
|
| 31 |
+
|
| 32 |
+
|
| 33 |
class YoloHeadTracker:
|
| 34 |
"""Lightweight head tracker using YOLO for face detection."""
|
| 35 |
|
|
|
|
| 45 |
|
| 46 |
try:
|
| 47 |
model_path = hf_hub_download(repo_id=model_repo, filename=model_filename)
|
| 48 |
+
self.model = cast(_YoloModel, YOLO(model_path).to(device))
|
| 49 |
logger.info("YOLO face detection model loaded from %s", model_repo)
|
| 50 |
except Exception as e:
|
| 51 |
logger.error("Failed to load YOLO model: %s", e)
|
src/reachy_mini_conversation_app/vision/head_tracking/yolo_process.py
CHANGED
|
@@ -137,6 +137,7 @@ class YoloHeadTrackerProcess:
|
|
| 137 |
self._messages: queue.Queue[tuple[str, object | None]] = queue.Queue()
|
| 138 |
self._next_request_id = 0
|
| 139 |
self._timed_out_request_id: int | None = None
|
|
|
|
| 140 |
self._tracker_name = "yolo"
|
| 141 |
|
| 142 |
module_path = "reachy_mini_conversation_app.vision.head_tracking.yolo_process"
|
|
@@ -304,8 +305,14 @@ class YoloHeadTrackerProcess:
|
|
| 304 |
request_id: int | None = None
|
| 305 |
try:
|
| 306 |
with self._send_lock:
|
| 307 |
-
#
|
| 308 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 309 |
return None, None
|
| 310 |
|
| 311 |
request_id = self._next_request_id
|
|
@@ -315,6 +322,7 @@ class YoloHeadTrackerProcess:
|
|
| 315 |
except TimeoutError as exc:
|
| 316 |
if request_id is not None:
|
| 317 |
self._timed_out_request_id = request_id
|
|
|
|
| 318 |
logger.warning("Head tracker %s communication timed out: %s", self._tracker_name, exc)
|
| 319 |
return None, None
|
| 320 |
except Exception as exc:
|
|
|
|
| 137 |
self._messages: queue.Queue[tuple[str, object | None]] = queue.Queue()
|
| 138 |
self._next_request_id = 0
|
| 139 |
self._timed_out_request_id: int | None = None
|
| 140 |
+
self._recovery_call_pending = False
|
| 141 |
self._tracker_name = "yolo"
|
| 142 |
|
| 143 |
module_path = "reachy_mini_conversation_app.vision.head_tracking.yolo_process"
|
|
|
|
| 305 |
request_id: int | None = None
|
| 306 |
try:
|
| 307 |
with self._send_lock:
|
| 308 |
+
# Reserve the immediate next call after a timeout for recovery
|
| 309 |
+
# only. Later calls may drain a delayed reply and continue with
|
| 310 |
+
# a fresh request in the same turn.
|
| 311 |
+
if self._recovery_call_pending:
|
| 312 |
+
self._recovery_call_pending = False
|
| 313 |
+
self._drain_timed_out_reply()
|
| 314 |
+
return None, None
|
| 315 |
+
if self._timed_out_request_id is not None and not self._drain_timed_out_reply():
|
| 316 |
return None, None
|
| 317 |
|
| 318 |
request_id = self._next_request_id
|
|
|
|
| 322 |
except TimeoutError as exc:
|
| 323 |
if request_id is not None:
|
| 324 |
self._timed_out_request_id = request_id
|
| 325 |
+
self._recovery_call_pending = True
|
| 326 |
logger.warning("Head tracker %s communication timed out: %s", self._tracker_name, exc)
|
| 327 |
return None, None
|
| 328 |
except Exception as exc:
|
src/reachy_mini_conversation_app/vision/local_vision.py
CHANGED
|
@@ -1,13 +1,16 @@
|
|
|
|
|
| 1 |
import os
|
| 2 |
import time
|
| 3 |
import logging
|
|
|
|
| 4 |
from dataclasses import dataclass
|
|
|
|
| 5 |
|
| 6 |
import numpy as np
|
| 7 |
import torch
|
| 8 |
from PIL import Image
|
| 9 |
from numpy.typing import NDArray
|
| 10 |
-
from transformers import AutoProcessor,
|
| 11 |
from huggingface_hub import snapshot_download
|
| 12 |
|
| 13 |
from reachy_mini_conversation_app.config import config
|
|
@@ -23,6 +26,48 @@ LOCAL_VISION_RESPONSE_INSTRUCTIONS = (
|
|
| 23 |
)
|
| 24 |
|
| 25 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
@dataclass
|
| 27 |
class VisionConfig:
|
| 28 |
"""Configuration for vision processing."""
|
|
@@ -41,8 +86,8 @@ class VisionProcessor:
|
|
| 41 |
"""Initialize the vision processor."""
|
| 42 |
self.vision_config = vision_config or VisionConfig()
|
| 43 |
self.device = self._determine_device()
|
| 44 |
-
self.processor:
|
| 45 |
-
self.model:
|
| 46 |
self._initialized = False
|
| 47 |
|
| 48 |
def _determine_device(self) -> str:
|
|
@@ -62,15 +107,21 @@ class VisionProcessor:
|
|
| 62 |
def initialize(self) -> None:
|
| 63 |
"""Load model and processor onto the selected device."""
|
| 64 |
logger.info("Loading SmolVLM2 model on %s (HF_HOME=%s)", self.device, config.HF_HOME)
|
| 65 |
-
processor
|
|
|
|
|
|
|
|
|
|
| 66 |
|
| 67 |
model_kwargs: dict[str, object] = {
|
| 68 |
"dtype": torch.bfloat16 if self.device == "cuda" else torch.float32,
|
| 69 |
}
|
| 70 |
|
| 71 |
-
model
|
| 72 |
-
|
| 73 |
-
|
|
|
|
|
|
|
|
|
|
| 74 |
)
|
| 75 |
model = model.to(self.device)
|
| 76 |
|
|
@@ -112,25 +163,21 @@ class VisionProcessor:
|
|
| 112 |
for attempt in range(self.vision_config.max_retries):
|
| 113 |
try:
|
| 114 |
inputs = processor.apply_chat_template(
|
| 115 |
-
messages,
|
| 116 |
add_generation_prompt=True,
|
| 117 |
tokenize=True,
|
| 118 |
return_dict=True,
|
| 119 |
return_tensors="pt",
|
| 120 |
)
|
| 121 |
-
inputs = inputs.to(self.device)
|
| 122 |
-
prompt_len =
|
| 123 |
-
input_ids = inputs.get("input_ids")
|
| 124 |
-
input_shape = getattr(input_ids, "shape", None)
|
| 125 |
-
if input_shape:
|
| 126 |
-
prompt_len = int(input_shape[-1])
|
| 127 |
|
| 128 |
with torch.inference_mode():
|
| 129 |
-
generated_ids = model.generate(
|
| 130 |
**inputs,
|
| 131 |
do_sample=False,
|
| 132 |
max_new_tokens=self.vision_config.max_new_tokens,
|
| 133 |
-
pad_token_id=processor.tokenizer.eos_token_id,
|
| 134 |
)
|
| 135 |
|
| 136 |
# Decode only the newly generated tokens, skipping the prompt
|
|
@@ -140,7 +187,7 @@ class VisionProcessor:
|
|
| 140 |
new_token_ids = generated_ids[:, prompt_len:]
|
| 141 |
else:
|
| 142 |
new_token_ids = [token_ids[prompt_len:] for token_ids in generated_ids]
|
| 143 |
-
response = processor.batch_decode(
|
| 144 |
new_token_ids,
|
| 145 |
skip_special_tokens=True,
|
| 146 |
)[0]
|
|
@@ -167,6 +214,14 @@ class VisionProcessor:
|
|
| 167 |
return f"Vision processing error after {self.vision_config.max_retries} attempts"
|
| 168 |
|
| 169 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 170 |
def initialize_vision_processor() -> VisionProcessor:
|
| 171 |
"""Download the vision model and return an initialized VisionProcessor."""
|
| 172 |
try:
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
import os
|
| 3 |
import time
|
| 4 |
import logging
|
| 5 |
+
from typing import Any, Protocol, cast
|
| 6 |
from dataclasses import dataclass
|
| 7 |
+
from collections.abc import Mapping, Sequence
|
| 8 |
|
| 9 |
import numpy as np
|
| 10 |
import torch
|
| 11 |
from PIL import Image
|
| 12 |
from numpy.typing import NDArray
|
| 13 |
+
from transformers import AutoProcessor, AutoModelForImageTextToText
|
| 14 |
from huggingface_hub import snapshot_download
|
| 15 |
|
| 16 |
from reachy_mini_conversation_app.config import config
|
|
|
|
| 26 |
)
|
| 27 |
|
| 28 |
|
| 29 |
+
class _VisionInputs(Mapping[str, object]):
|
| 30 |
+
"""Tokenized processor inputs that can be moved to the inference device."""
|
| 31 |
+
|
| 32 |
+
def to(self, device: str) -> _VisionInputs:
|
| 33 |
+
"""Move inputs to the selected inference device."""
|
| 34 |
+
raise NotImplementedError
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
class _VisionTokenizer(Protocol):
|
| 38 |
+
"""Tokenizer attributes used by generation."""
|
| 39 |
+
|
| 40 |
+
eos_token_id: int | None
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
class _VisionProcessor(Protocol):
|
| 44 |
+
"""Small interface required from the Hugging Face image-text processor."""
|
| 45 |
+
|
| 46 |
+
tokenizer: _VisionTokenizer
|
| 47 |
+
|
| 48 |
+
def apply_chat_template(
|
| 49 |
+
self,
|
| 50 |
+
conversation: object,
|
| 51 |
+
*,
|
| 52 |
+
add_generation_prompt: bool,
|
| 53 |
+
tokenize: bool,
|
| 54 |
+
return_dict: bool,
|
| 55 |
+
return_tensors: str,
|
| 56 |
+
) -> _VisionInputs: ...
|
| 57 |
+
|
| 58 |
+
def batch_decode(self, sequences: object, *, skip_special_tokens: bool) -> list[str]: ...
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
class _VisionModel(Protocol):
|
| 62 |
+
"""Small interface required from the Hugging Face image-text model."""
|
| 63 |
+
|
| 64 |
+
def to(self, device: str) -> _VisionModel: ...
|
| 65 |
+
|
| 66 |
+
def eval(self) -> object: ...
|
| 67 |
+
|
| 68 |
+
def generate(self, **kwargs: object) -> Any: ...
|
| 69 |
+
|
| 70 |
+
|
| 71 |
@dataclass
|
| 72 |
class VisionConfig:
|
| 73 |
"""Configuration for vision processing."""
|
|
|
|
| 86 |
"""Initialize the vision processor."""
|
| 87 |
self.vision_config = vision_config or VisionConfig()
|
| 88 |
self.device = self._determine_device()
|
| 89 |
+
self.processor: _VisionProcessor | None = None
|
| 90 |
+
self.model: _VisionModel | None = None
|
| 91 |
self._initialized = False
|
| 92 |
|
| 93 |
def _determine_device(self) -> str:
|
|
|
|
| 107 |
def initialize(self) -> None:
|
| 108 |
"""Load model and processor onto the selected device."""
|
| 109 |
logger.info("Loading SmolVLM2 model on %s (HF_HOME=%s)", self.device, config.HF_HOME)
|
| 110 |
+
processor = cast(
|
| 111 |
+
_VisionProcessor,
|
| 112 |
+
AutoProcessor.from_pretrained(self.vision_config.model_path), # type: ignore[no-untyped-call]
|
| 113 |
+
)
|
| 114 |
|
| 115 |
model_kwargs: dict[str, object] = {
|
| 116 |
"dtype": torch.bfloat16 if self.device == "cuda" else torch.float32,
|
| 117 |
}
|
| 118 |
|
| 119 |
+
model = cast(
|
| 120 |
+
_VisionModel,
|
| 121 |
+
AutoModelForImageTextToText.from_pretrained(
|
| 122 |
+
self.vision_config.model_path,
|
| 123 |
+
**model_kwargs,
|
| 124 |
+
),
|
| 125 |
)
|
| 126 |
model = model.to(self.device)
|
| 127 |
|
|
|
|
| 163 |
for attempt in range(self.vision_config.max_retries):
|
| 164 |
try:
|
| 165 |
inputs = processor.apply_chat_template(
|
| 166 |
+
messages,
|
| 167 |
add_generation_prompt=True,
|
| 168 |
tokenize=True,
|
| 169 |
return_dict=True,
|
| 170 |
return_tensors="pt",
|
| 171 |
)
|
| 172 |
+
inputs = inputs.to(self.device)
|
| 173 |
+
prompt_len = _last_shape_dim(inputs.get("input_ids"))
|
|
|
|
|
|
|
|
|
|
|
|
|
| 174 |
|
| 175 |
with torch.inference_mode():
|
| 176 |
+
generated_ids = model.generate(
|
| 177 |
**inputs,
|
| 178 |
do_sample=False,
|
| 179 |
max_new_tokens=self.vision_config.max_new_tokens,
|
| 180 |
+
pad_token_id=processor.tokenizer.eos_token_id,
|
| 181 |
)
|
| 182 |
|
| 183 |
# Decode only the newly generated tokens, skipping the prompt
|
|
|
|
| 187 |
new_token_ids = generated_ids[:, prompt_len:]
|
| 188 |
else:
|
| 189 |
new_token_ids = [token_ids[prompt_len:] for token_ids in generated_ids]
|
| 190 |
+
response = processor.batch_decode(
|
| 191 |
new_token_ids,
|
| 192 |
skip_special_tokens=True,
|
| 193 |
)[0]
|
|
|
|
| 214 |
return f"Vision processing error after {self.vision_config.max_retries} attempts"
|
| 215 |
|
| 216 |
|
| 217 |
+
def _last_shape_dim(value: object) -> int | None:
|
| 218 |
+
"""Return the last dimension from tensor-like objects that expose a shape."""
|
| 219 |
+
shape = getattr(value, "shape", None)
|
| 220 |
+
if not isinstance(shape, Sequence) or not shape:
|
| 221 |
+
return None
|
| 222 |
+
return int(shape[-1])
|
| 223 |
+
|
| 224 |
+
|
| 225 |
def initialize_vision_processor() -> VisionProcessor:
|
| 226 |
"""Download the vision model and return an initialized VisionProcessor."""
|
| 227 |
try:
|
tests/audio/test_head_wobbler.py
CHANGED
|
@@ -8,6 +8,7 @@ from typing import Any, List, Tuple
|
|
| 8 |
from collections.abc import Callable
|
| 9 |
|
| 10 |
import numpy as np
|
|
|
|
| 11 |
|
| 12 |
from reachy_mini_conversation_app.audio.head_wobbler import HeadWobbler
|
| 13 |
|
|
@@ -18,7 +19,11 @@ def _make_audio_chunk(duration_s: float = 0.3, frequency_hz: float = 220.0) -> s
|
|
| 18 |
return base64.b64encode(pcm.tobytes()).decode("ascii")
|
| 19 |
|
| 20 |
|
| 21 |
-
def _make_pcm(
|
|
|
|
|
|
|
|
|
|
|
|
|
| 22 |
"""Generate a mono PCM16 sine wave at the requested sample rate."""
|
| 23 |
sample_count = int(sample_rate * duration_s)
|
| 24 |
t = np.linspace(0, duration_s, sample_count, endpoint=False)
|
|
|
|
| 8 |
from collections.abc import Callable
|
| 9 |
|
| 10 |
import numpy as np
|
| 11 |
+
from numpy.typing import NDArray
|
| 12 |
|
| 13 |
from reachy_mini_conversation_app.audio.head_wobbler import HeadWobbler
|
| 14 |
|
|
|
|
| 19 |
return base64.b64encode(pcm.tobytes()).decode("ascii")
|
| 20 |
|
| 21 |
|
| 22 |
+
def _make_pcm(
|
| 23 |
+
duration_s: float = 0.3,
|
| 24 |
+
frequency_hz: float = 220.0,
|
| 25 |
+
sample_rate: int = 24000,
|
| 26 |
+
) -> NDArray[np.int16]:
|
| 27 |
"""Generate a mono PCM16 sine wave at the requested sample rate."""
|
| 28 |
sample_count = int(sample_rate * duration_s)
|
| 29 |
t = np.linspace(0, duration_s, sample_count, endpoint=False)
|
tests/audio/test_startup_config.py
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Tests for Reachy Mini audio startup configuration."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
from types import SimpleNamespace
|
| 5 |
+
|
| 6 |
+
from reachy_mini_conversation_app.audio.startup_config import (
|
| 7 |
+
AUDIO_STARTUP_CONFIG,
|
| 8 |
+
WRITE_SETTLE_SECONDS,
|
| 9 |
+
apply_audio_startup_config,
|
| 10 |
+
)
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class FakeAudio:
|
| 14 |
+
"""Fake SDK audio wrapper."""
|
| 15 |
+
|
| 16 |
+
def __init__(self, *, result: bool = True, error: Exception | None = None) -> None:
|
| 17 |
+
"""Initialize the fake audio wrapper."""
|
| 18 |
+
self.result = result
|
| 19 |
+
self.error = error
|
| 20 |
+
self.calls: list[tuple[object, bool, float]] = []
|
| 21 |
+
|
| 22 |
+
def apply_audio_config(
|
| 23 |
+
self,
|
| 24 |
+
config: object,
|
| 25 |
+
*,
|
| 26 |
+
verify: bool = True,
|
| 27 |
+
write_settle_seconds: float = WRITE_SETTLE_SECONDS,
|
| 28 |
+
) -> bool:
|
| 29 |
+
"""Record SDK audio config calls."""
|
| 30 |
+
if self.error is not None:
|
| 31 |
+
raise self.error
|
| 32 |
+
self.calls.append((config, verify, write_settle_seconds))
|
| 33 |
+
return self.result
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def test_apply_audio_startup_config_uses_sdk_audio_config_api() -> None:
|
| 37 |
+
"""Startup config should delegate writes and verification to the SDK audio API."""
|
| 38 |
+
audio = FakeAudio()
|
| 39 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=audio))
|
| 40 |
+
|
| 41 |
+
applied = apply_audio_startup_config(robot)
|
| 42 |
+
|
| 43 |
+
assert applied is True
|
| 44 |
+
assert audio.calls == [(AUDIO_STARTUP_CONFIG, True, WRITE_SETTLE_SECONDS)]
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def test_apply_audio_startup_config_forwards_sdk_options() -> None:
|
| 48 |
+
"""SDK verification options should stay configurable for tests and callers."""
|
| 49 |
+
audio = FakeAudio()
|
| 50 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=audio))
|
| 51 |
+
|
| 52 |
+
applied = apply_audio_startup_config(robot, verify=False, write_settle_seconds=0)
|
| 53 |
+
|
| 54 |
+
assert applied is True
|
| 55 |
+
assert audio.calls == [(AUDIO_STARTUP_CONFIG, False, 0)]
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def test_apply_audio_startup_config_returns_false_without_audio() -> None:
|
| 59 |
+
"""Startup should continue when the SDK audio object is unavailable."""
|
| 60 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=None))
|
| 61 |
+
|
| 62 |
+
applied = apply_audio_startup_config(robot)
|
| 63 |
+
|
| 64 |
+
assert applied is False
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
def test_apply_audio_startup_config_returns_false_without_sdk_api() -> None:
|
| 68 |
+
"""Startup should continue when the installed SDK does not expose audio config helpers."""
|
| 69 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=object()))
|
| 70 |
+
|
| 71 |
+
applied = apply_audio_startup_config(robot)
|
| 72 |
+
|
| 73 |
+
assert applied is False
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
def test_apply_audio_startup_config_returns_false_when_sdk_returns_false() -> None:
|
| 77 |
+
"""SDK application failures should be reported without raising."""
|
| 78 |
+
audio = FakeAudio(result=False)
|
| 79 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=audio))
|
| 80 |
+
|
| 81 |
+
applied = apply_audio_startup_config(robot)
|
| 82 |
+
|
| 83 |
+
assert applied is False
|
| 84 |
+
assert audio.calls == [(AUDIO_STARTUP_CONFIG, True, WRITE_SETTLE_SECONDS)]
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def test_apply_audio_startup_config_returns_false_when_sdk_raises() -> None:
|
| 88 |
+
"""Unexpected SDK audio config errors should not prevent app startup."""
|
| 89 |
+
audio = FakeAudio(error=RuntimeError("audio board unavailable"))
|
| 90 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=audio))
|
| 91 |
+
|
| 92 |
+
applied = apply_audio_startup_config(robot)
|
| 93 |
+
|
| 94 |
+
assert applied is False
|
| 95 |
+
assert audio.calls == []
|
tests/test_config_name_collisions.py
CHANGED
|
@@ -59,3 +59,86 @@ def test_config_raises_when_selected_external_profile_is_missing(
|
|
| 59 |
|
| 60 |
with pytest.raises(RuntimeError, match="Selected profile 'missing_profile' was not found"):
|
| 61 |
config_mod.Config()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 59 |
|
| 60 |
with pytest.raises(RuntimeError, match="Selected profile 'missing_profile' was not found"):
|
| 61 |
config_mod.Config()
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def test_backend_provider_defaults_to_hf_when_unset() -> None:
|
| 65 |
+
"""Non-Gemini models should default to the Hugging Face backend."""
|
| 66 |
+
assert config_mod._normalize_backend_provider(None, None) == config_mod.HF_BACKEND
|
| 67 |
+
assert config_mod._normalize_backend_provider("", None) == config_mod.HF_BACKEND
|
| 68 |
+
assert config_mod._normalize_backend_provider(None, "gpt-realtime") == config_mod.HF_BACKEND
|
| 69 |
+
assert config_mod._normalize_backend_provider(None, "gemini-3.1-flash-live-preview") == config_mod.GEMINI_BACKEND
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def test_backend_provider_rejects_explicit_unknown_backend() -> None:
|
| 73 |
+
"""An explicit backend typo should fail instead of falling through to the default backend."""
|
| 74 |
+
with pytest.raises(ValueError, match="Invalid BACKEND_PROVIDER='openia'"):
|
| 75 |
+
config_mod._normalize_backend_provider("openia", None)
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def test_huggingface_backend_does_not_resolve_model_name() -> None:
|
| 79 |
+
"""Hugging Face should rely on the server's model selection."""
|
| 80 |
+
assert config_mod._resolve_model_name(config_mod.HF_BACKEND, None) == ""
|
| 81 |
+
assert config_mod._resolve_model_name(config_mod.HF_BACKEND, "gpt-realtime") == ""
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def test_hf_default_session_url_uses_stable_space_proxy() -> None:
|
| 85 |
+
"""The app should not embed the raw, replaceable Inference Endpoint allocator URL."""
|
| 86 |
+
assert config_mod.HF_DEFAULTS.session_url == "https://pollen-robotics-reachy-mini-realtime-url.hf.space/session"
|
| 87 |
+
assert ".aws.endpoints.huggingface.cloud" not in config_mod.HF_DEFAULTS.session_url
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
def test_refresh_runtime_config_reloads_hf_runtime_fields(monkeypatch: pytest.MonkeyPatch) -> None:
|
| 91 |
+
"""Instance-local .env reloads should update every env-backed Hugging Face runtime field."""
|
| 92 |
+
monkeypatch.setenv("HF_TOKEN", "hf-runtime-token")
|
| 93 |
+
monkeypatch.setenv("HF_HOME", "/tmp/reachy-hf-cache")
|
| 94 |
+
monkeypatch.setenv("LOCAL_VISION_MODEL", "test/local-vision-model")
|
| 95 |
+
|
| 96 |
+
monkeypatch.setattr(config_mod.config, "HF_TOKEN", None)
|
| 97 |
+
monkeypatch.setattr(config_mod.config, "HF_HOME", "./old-cache")
|
| 98 |
+
monkeypatch.setattr(config_mod.config, "LOCAL_VISION_MODEL", "old/model")
|
| 99 |
+
|
| 100 |
+
config_mod.refresh_runtime_config_from_env()
|
| 101 |
+
|
| 102 |
+
assert config_mod.config.HF_TOKEN == "hf-runtime-token"
|
| 103 |
+
assert config_mod.config.HF_HOME == "/tmp/reachy-hf-cache"
|
| 104 |
+
assert config_mod.config.LOCAL_VISION_MODEL == "test/local-vision-model"
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
@pytest.mark.parametrize(
|
| 108 |
+
("configured_mode", "session_url", "direct_ws_url", "expected_mode", "expected_has_target"),
|
| 109 |
+
[
|
| 110 |
+
("local", "https://hf.example.test/session", None, "local", False),
|
| 111 |
+
("deployed", "https://hf.example.test/session", "ws://127.0.0.1:8765/v1/realtime", "deployed", True),
|
| 112 |
+
("local", None, "ws://127.0.0.1:8765/v1/realtime", "local", True),
|
| 113 |
+
("deployed", None, "ws://127.0.0.1:8765/v1/realtime", "deployed", False),
|
| 114 |
+
],
|
| 115 |
+
)
|
| 116 |
+
def test_hf_connection_selection_uses_explicit_mode_for_target(
|
| 117 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 118 |
+
configured_mode: str | None,
|
| 119 |
+
session_url: str | None,
|
| 120 |
+
direct_ws_url: str | None,
|
| 121 |
+
expected_mode: str,
|
| 122 |
+
expected_has_target: bool,
|
| 123 |
+
) -> None:
|
| 124 |
+
"""Hugging Face selection should use the configured mode without inferring from URLs."""
|
| 125 |
+
monkeypatch.setattr(config_mod.config, "HF_REALTIME_CONNECTION_MODE", configured_mode)
|
| 126 |
+
monkeypatch.setattr(config_mod.config, "HF_REALTIME_SESSION_URL", session_url)
|
| 127 |
+
monkeypatch.setattr(config_mod.config, "HF_REALTIME_WS_URL", direct_ws_url)
|
| 128 |
+
|
| 129 |
+
selection = config_mod.get_hf_connection_selection()
|
| 130 |
+
|
| 131 |
+
assert selection.mode == expected_mode
|
| 132 |
+
assert selection.has_target is expected_has_target
|
| 133 |
+
assert selection.session_url == session_url
|
| 134 |
+
assert selection.direct_ws_url == direct_ws_url
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
def test_hf_connection_selection_requires_mode(monkeypatch: pytest.MonkeyPatch) -> None:
|
| 138 |
+
"""Hugging Face selection should fail instead of inferring a missing mode."""
|
| 139 |
+
monkeypatch.setattr(config_mod.config, "HF_REALTIME_CONNECTION_MODE", None)
|
| 140 |
+
monkeypatch.setattr(config_mod.config, "HF_REALTIME_SESSION_URL", "https://hf.example.test/session")
|
| 141 |
+
monkeypatch.setattr(config_mod.config, "HF_REALTIME_WS_URL", "ws://127.0.0.1:8765/v1/realtime")
|
| 142 |
+
|
| 143 |
+
with pytest.raises(RuntimeError, match="HF_REALTIME_CONNECTION_MODE must be set"):
|
| 144 |
+
config_mod.get_hf_connection_selection()
|
tests/test_console.py
CHANGED
|
@@ -1,18 +1,26 @@
|
|
| 1 |
"""Tests for the headless console stream."""
|
| 2 |
|
|
|
|
| 3 |
import asyncio
|
| 4 |
import threading
|
| 5 |
from types import SimpleNamespace
|
|
|
|
|
|
|
| 6 |
from unittest.mock import AsyncMock, MagicMock
|
| 7 |
|
| 8 |
import numpy as np
|
| 9 |
import pytest
|
| 10 |
from fastapi import FastAPI
|
|
|
|
| 11 |
from fastapi.testclient import TestClient
|
| 12 |
|
| 13 |
from reachy_mini.media.media_manager import MediaBackend
|
| 14 |
from reachy_mini_conversation_app.config import GEMINI_AVAILABLE_VOICES, config
|
| 15 |
-
from reachy_mini_conversation_app.console import LocalStream
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
from reachy_mini_conversation_app.headless_personality_ui import mount_personality_routes
|
| 17 |
|
| 18 |
|
|
@@ -23,7 +31,7 @@ def test_clear_audio_queue_prefers_clear_player_when_available() -> None:
|
|
| 23 |
clear_player=MagicMock(),
|
| 24 |
clear_output_buffer=MagicMock(),
|
| 25 |
)
|
| 26 |
-
robot = SimpleNamespace(media=SimpleNamespace(audio=audio, backend=
|
| 27 |
stream = LocalStream(handler, robot)
|
| 28 |
|
| 29 |
stream.clear_audio_queue()
|
|
@@ -75,10 +83,10 @@ async def test_play_loop_feeds_head_wobbler_with_local_playback_delay() -> None:
|
|
| 75 |
class Handler:
|
| 76 |
def __init__(self) -> None:
|
| 77 |
self.deps = SimpleNamespace(head_wobbler=head_wobbler)
|
| 78 |
-
self.output_queue = asyncio.Queue()
|
| 79 |
self._emitted = False
|
| 80 |
|
| 81 |
-
async def emit(self):
|
| 82 |
if not self._emitted:
|
| 83 |
self._emitted = True
|
| 84 |
return (24000, chunk.copy())
|
|
@@ -90,7 +98,7 @@ async def test_play_loop_feeds_head_wobbler_with_local_playback_delay() -> None:
|
|
| 90 |
)
|
| 91 |
media = SimpleNamespace(
|
| 92 |
audio=audio,
|
| 93 |
-
backend=
|
| 94 |
get_output_audio_samplerate=lambda: 24000,
|
| 95 |
push_audio_sample=MagicMock(),
|
| 96 |
)
|
|
@@ -117,8 +125,8 @@ async def test_play_loop_feeds_head_wobbler_with_local_playback_delay() -> None:
|
|
| 117 |
|
| 118 |
|
| 119 |
def test_backend_config_persists_gemini_selection_and_status(
|
| 120 |
-
tmp_path,
|
| 121 |
-
monkeypatch,
|
| 122 |
) -> None:
|
| 123 |
"""Settings API should persist Gemini backend choice and token."""
|
| 124 |
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
|
@@ -171,8 +179,8 @@ def test_backend_config_persists_gemini_selection_and_status(
|
|
| 171 |
|
| 172 |
|
| 173 |
def test_backend_config_preserves_explicit_model_override_when_saving_key(
|
| 174 |
-
tmp_path,
|
| 175 |
-
monkeypatch,
|
| 176 |
) -> None:
|
| 177 |
"""Saving credentials should not reset a custom model override."""
|
| 178 |
custom_model = "gpt-4o-realtime-preview-2025-06-03"
|
|
@@ -211,7 +219,219 @@ def test_backend_config_preserves_explicit_model_override_when_saving_key(
|
|
| 211 |
assert "OPENAI_API_KEY=openai-test-key" in env_text
|
| 212 |
|
| 213 |
|
| 214 |
-
def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 215 |
"""Headless personality UI should expose Gemini voices when Gemini is selected."""
|
| 216 |
monkeypatch.setattr(config, "BACKEND_PROVIDER", "gemini")
|
| 217 |
monkeypatch.setattr(config, "MODEL_NAME", "gemini-3.1-flash-live-preview")
|
|
@@ -227,6 +447,22 @@ def test_headless_personality_routes_return_gemini_voices_when_backend_selected(
|
|
| 227 |
assert response.json() == GEMINI_AVAILABLE_VOICES
|
| 228 |
|
| 229 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 230 |
def test_headless_personality_routes_apply_voice_accepts_query_param() -> None:
|
| 231 |
"""Headless personality UI should apply a voice change from a POST query param."""
|
| 232 |
app = FastAPI()
|
|
@@ -258,3 +494,116 @@ def test_headless_personality_routes_apply_voice_accepts_query_param() -> None:
|
|
| 258 |
loop.call_soon_threadsafe(loop.stop)
|
| 259 |
thread.join(timeout=1.0)
|
| 260 |
loop.close()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
"""Tests for the headless console stream."""
|
| 2 |
|
| 3 |
+
import sys
|
| 4 |
import asyncio
|
| 5 |
import threading
|
| 6 |
from types import SimpleNamespace
|
| 7 |
+
from typing import Any
|
| 8 |
+
from pathlib import Path
|
| 9 |
from unittest.mock import AsyncMock, MagicMock
|
| 10 |
|
| 11 |
import numpy as np
|
| 12 |
import pytest
|
| 13 |
from fastapi import FastAPI
|
| 14 |
+
from numpy.typing import NDArray
|
| 15 |
from fastapi.testclient import TestClient
|
| 16 |
|
| 17 |
from reachy_mini.media.media_manager import MediaBackend
|
| 18 |
from reachy_mini_conversation_app.config import GEMINI_AVAILABLE_VOICES, config
|
| 19 |
+
from reachy_mini_conversation_app.console import LOCAL_PLAYER_BACKEND, LocalStream
|
| 20 |
+
from reachy_mini_conversation_app.startup_settings import (
|
| 21 |
+
StartupSettings,
|
| 22 |
+
load_startup_settings_into_runtime,
|
| 23 |
+
)
|
| 24 |
from reachy_mini_conversation_app.headless_personality_ui import mount_personality_routes
|
| 25 |
|
| 26 |
|
|
|
|
| 31 |
clear_player=MagicMock(),
|
| 32 |
clear_output_buffer=MagicMock(),
|
| 33 |
)
|
| 34 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=audio, backend=LOCAL_PLAYER_BACKEND))
|
| 35 |
stream = LocalStream(handler, robot)
|
| 36 |
|
| 37 |
stream.clear_audio_queue()
|
|
|
|
| 83 |
class Handler:
|
| 84 |
def __init__(self) -> None:
|
| 85 |
self.deps = SimpleNamespace(head_wobbler=head_wobbler)
|
| 86 |
+
self.output_queue: asyncio.Queue[Any] = asyncio.Queue()
|
| 87 |
self._emitted = False
|
| 88 |
|
| 89 |
+
async def emit(self) -> tuple[int, NDArray[np.int16]] | None:
|
| 90 |
if not self._emitted:
|
| 91 |
self._emitted = True
|
| 92 |
return (24000, chunk.copy())
|
|
|
|
| 98 |
)
|
| 99 |
media = SimpleNamespace(
|
| 100 |
audio=audio,
|
| 101 |
+
backend=LOCAL_PLAYER_BACKEND,
|
| 102 |
get_output_audio_samplerate=lambda: 24000,
|
| 103 |
push_audio_sample=MagicMock(),
|
| 104 |
)
|
|
|
|
| 125 |
|
| 126 |
|
| 127 |
def test_backend_config_persists_gemini_selection_and_status(
|
| 128 |
+
tmp_path: Path,
|
| 129 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 130 |
) -> None:
|
| 131 |
"""Settings API should persist Gemini backend choice and token."""
|
| 132 |
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
|
|
|
| 179 |
|
| 180 |
|
| 181 |
def test_backend_config_preserves_explicit_model_override_when_saving_key(
|
| 182 |
+
tmp_path: Path,
|
| 183 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 184 |
) -> None:
|
| 185 |
"""Saving credentials should not reset a custom model override."""
|
| 186 |
custom_model = "gpt-4o-realtime-preview-2025-06-03"
|
|
|
|
| 219 |
assert "OPENAI_API_KEY=openai-test-key" in env_text
|
| 220 |
|
| 221 |
|
| 222 |
+
def test_backend_config_persists_local_hf_selection_and_status(
|
| 223 |
+
tmp_path: Path,
|
| 224 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 225 |
+
) -> None:
|
| 226 |
+
"""Settings API should persist a direct Hugging Face websocket target."""
|
| 227 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 228 |
+
monkeypatch.setattr(config, "MODEL_NAME", "gpt-realtime")
|
| 229 |
+
monkeypatch.setattr(config, "HF_REALTIME_CONNECTION_MODE", "deployed")
|
| 230 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", None)
|
| 231 |
+
monkeypatch.setattr(config, "HF_REALTIME_WS_URL", None)
|
| 232 |
+
monkeypatch.setenv("BACKEND_PROVIDER", "openai")
|
| 233 |
+
monkeypatch.setenv("MODEL_NAME", "gpt-realtime")
|
| 234 |
+
monkeypatch.delenv("HF_REALTIME_CONNECTION_MODE", raising=False)
|
| 235 |
+
monkeypatch.delenv("HF_REALTIME_SESSION_URL", raising=False)
|
| 236 |
+
monkeypatch.delenv("HF_REALTIME_WS_URL", raising=False)
|
| 237 |
+
|
| 238 |
+
app = FastAPI()
|
| 239 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=None, backend=None))
|
| 240 |
+
stream = LocalStream(MagicMock(), robot, settings_app=app, instance_path=str(tmp_path))
|
| 241 |
+
stream._init_settings_ui_if_needed()
|
| 242 |
+
|
| 243 |
+
client = TestClient(app)
|
| 244 |
+
response = client.post(
|
| 245 |
+
"/backend_config",
|
| 246 |
+
json={
|
| 247 |
+
"backend": "huggingface",
|
| 248 |
+
"hf_mode": "local",
|
| 249 |
+
"hf_host": "localhost",
|
| 250 |
+
"hf_port": 8765,
|
| 251 |
+
},
|
| 252 |
+
)
|
| 253 |
+
|
| 254 |
+
assert response.status_code == 200
|
| 255 |
+
data = response.json()
|
| 256 |
+
assert data["ok"] is True
|
| 257 |
+
assert data["backend_provider"] == "huggingface"
|
| 258 |
+
assert data["active_backend"] == "openai"
|
| 259 |
+
assert data["has_hf_ws_url"] is True
|
| 260 |
+
assert data["has_hf_connection"] is True
|
| 261 |
+
assert data["hf_connection_mode"] == "local"
|
| 262 |
+
assert data["hf_direct_host"] == "localhost"
|
| 263 |
+
assert data["hf_direct_port"] == 8765
|
| 264 |
+
assert data["requires_restart"] is True
|
| 265 |
+
|
| 266 |
+
env_text = (tmp_path / ".env").read_text(encoding="utf-8")
|
| 267 |
+
env_lines = env_text.splitlines()
|
| 268 |
+
assert "BACKEND_PROVIDER=huggingface" in env_text
|
| 269 |
+
assert "HF_REALTIME_CONNECTION_MODE=local" in env_text
|
| 270 |
+
assert "HF_REALTIME_WS_URL=ws://localhost:8765/v1/realtime" in env_text
|
| 271 |
+
assert not any(line.startswith("MODEL_NAME=") for line in env_lines)
|
| 272 |
+
|
| 273 |
+
|
| 274 |
+
def test_backend_config_persists_deployed_mode_without_clearing_local_hf_ws_url(
|
| 275 |
+
tmp_path: Path,
|
| 276 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 277 |
+
) -> None:
|
| 278 |
+
"""Saving deployed mode should make env selection explicit and remove stale allocator URLs."""
|
| 279 |
+
env_path = tmp_path / ".env"
|
| 280 |
+
env_path.write_text(
|
| 281 |
+
"BACKEND_PROVIDER=huggingface\n"
|
| 282 |
+
"HF_REALTIME_SESSION_URL=https://lb.example.test/session\n"
|
| 283 |
+
"HF_REALTIME_WS_URL=ws://localhost:8765/v1/realtime\n",
|
| 284 |
+
encoding="utf-8",
|
| 285 |
+
)
|
| 286 |
+
|
| 287 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 288 |
+
monkeypatch.setattr(config, "MODEL_NAME", "gpt-realtime")
|
| 289 |
+
monkeypatch.setattr(config, "HF_REALTIME_CONNECTION_MODE", "deployed")
|
| 290 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", "https://lb.example.test/session")
|
| 291 |
+
monkeypatch.setattr(config, "HF_REALTIME_WS_URL", "ws://localhost:8765/v1/realtime")
|
| 292 |
+
monkeypatch.setenv("BACKEND_PROVIDER", "huggingface")
|
| 293 |
+
monkeypatch.setenv("MODEL_NAME", "gpt-realtime")
|
| 294 |
+
monkeypatch.delenv("HF_REALTIME_CONNECTION_MODE", raising=False)
|
| 295 |
+
monkeypatch.setenv("HF_REALTIME_SESSION_URL", "https://lb.example.test/session")
|
| 296 |
+
monkeypatch.setenv("HF_REALTIME_WS_URL", "ws://localhost:8765/v1/realtime")
|
| 297 |
+
|
| 298 |
+
app = FastAPI()
|
| 299 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=None, backend=None))
|
| 300 |
+
stream = LocalStream(MagicMock(), robot, settings_app=app, instance_path=str(tmp_path))
|
| 301 |
+
stream._init_settings_ui_if_needed()
|
| 302 |
+
|
| 303 |
+
client = TestClient(app)
|
| 304 |
+
response = client.post(
|
| 305 |
+
"/backend_config",
|
| 306 |
+
json={
|
| 307 |
+
"backend": "huggingface",
|
| 308 |
+
"hf_mode": "deployed",
|
| 309 |
+
},
|
| 310 |
+
)
|
| 311 |
+
|
| 312 |
+
assert response.status_code == 200
|
| 313 |
+
data = response.json()
|
| 314 |
+
assert data["ok"] is True
|
| 315 |
+
assert data["has_hf_session_url"] is True
|
| 316 |
+
assert data["has_hf_ws_url"] is True
|
| 317 |
+
assert data["hf_connection_mode"] == "deployed"
|
| 318 |
+
|
| 319 |
+
env_text = env_path.read_text(encoding="utf-8")
|
| 320 |
+
assert "HF_REALTIME_CONNECTION_MODE=deployed" in env_text
|
| 321 |
+
assert "HF_REALTIME_SESSION_URL=" not in env_text
|
| 322 |
+
assert "HF_REALTIME_WS_URL=ws://localhost:8765/v1/realtime" in env_text
|
| 323 |
+
|
| 324 |
+
|
| 325 |
+
def test_backend_config_switches_to_saved_local_hf_connection_without_payload_target(
|
| 326 |
+
tmp_path: Path,
|
| 327 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 328 |
+
) -> None:
|
| 329 |
+
"""Switching back to a saved local Hugging Face backend should reuse the persisted target."""
|
| 330 |
+
env_path = tmp_path / ".env"
|
| 331 |
+
env_path.write_text(
|
| 332 |
+
"BACKEND_PROVIDER=openai\n"
|
| 333 |
+
"MODEL_NAME=gpt-realtime\n"
|
| 334 |
+
"HF_REALTIME_CONNECTION_MODE=local\n"
|
| 335 |
+
"HF_REALTIME_WS_URL=ws://192.168.1.42:8766/v1/realtime\n",
|
| 336 |
+
encoding="utf-8",
|
| 337 |
+
)
|
| 338 |
+
|
| 339 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 340 |
+
monkeypatch.setattr(config, "MODEL_NAME", "gpt-realtime")
|
| 341 |
+
monkeypatch.setattr(config, "HF_REALTIME_CONNECTION_MODE", "local")
|
| 342 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", None)
|
| 343 |
+
monkeypatch.setattr(config, "HF_REALTIME_WS_URL", "ws://192.168.1.42:8766/v1/realtime")
|
| 344 |
+
monkeypatch.setenv("BACKEND_PROVIDER", "openai")
|
| 345 |
+
monkeypatch.setenv("MODEL_NAME", "gpt-realtime")
|
| 346 |
+
monkeypatch.setenv("HF_REALTIME_CONNECTION_MODE", "local")
|
| 347 |
+
monkeypatch.setenv("HF_REALTIME_WS_URL", "ws://192.168.1.42:8766/v1/realtime")
|
| 348 |
+
|
| 349 |
+
app = FastAPI()
|
| 350 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=None, backend=None))
|
| 351 |
+
stream = LocalStream(MagicMock(), robot, settings_app=app, instance_path=str(tmp_path))
|
| 352 |
+
stream._init_settings_ui_if_needed()
|
| 353 |
+
|
| 354 |
+
client = TestClient(app)
|
| 355 |
+
response = client.post(
|
| 356 |
+
"/backend_config",
|
| 357 |
+
json={"backend": "huggingface"},
|
| 358 |
+
)
|
| 359 |
+
|
| 360 |
+
assert response.status_code == 200
|
| 361 |
+
data = response.json()
|
| 362 |
+
assert data["ok"] is True
|
| 363 |
+
assert data["backend_provider"] == "huggingface"
|
| 364 |
+
assert data["hf_connection_mode"] == "local"
|
| 365 |
+
assert data["hf_direct_host"] == "192.168.1.42"
|
| 366 |
+
assert data["hf_direct_port"] == 8766
|
| 367 |
+
|
| 368 |
+
env_text = env_path.read_text(encoding="utf-8")
|
| 369 |
+
assert "BACKEND_PROVIDER=huggingface" in env_text
|
| 370 |
+
assert "HF_REALTIME_CONNECTION_MODE=local" in env_text
|
| 371 |
+
assert "HF_REALTIME_WS_URL=ws://192.168.1.42:8766/v1/realtime" in env_text
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
def test_backend_config_rejects_invalid_hf_port_zero(
|
| 375 |
+
tmp_path: Path,
|
| 376 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 377 |
+
) -> None:
|
| 378 |
+
"""Settings API should reject invalid local Hugging Face ports from direct callers."""
|
| 379 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 380 |
+
monkeypatch.setattr(config, "HF_REALTIME_CONNECTION_MODE", "deployed")
|
| 381 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", None)
|
| 382 |
+
monkeypatch.setattr(config, "HF_REALTIME_WS_URL", None)
|
| 383 |
+
|
| 384 |
+
app = FastAPI()
|
| 385 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=None, backend=None))
|
| 386 |
+
stream = LocalStream(MagicMock(), robot, settings_app=app, instance_path=str(tmp_path))
|
| 387 |
+
stream._init_settings_ui_if_needed()
|
| 388 |
+
|
| 389 |
+
client = TestClient(app)
|
| 390 |
+
response = client.post(
|
| 391 |
+
"/backend_config",
|
| 392 |
+
json={
|
| 393 |
+
"backend": "huggingface",
|
| 394 |
+
"hf_mode": "local",
|
| 395 |
+
"hf_host": "localhost",
|
| 396 |
+
"hf_port": 0,
|
| 397 |
+
},
|
| 398 |
+
)
|
| 399 |
+
|
| 400 |
+
assert response.status_code == 400
|
| 401 |
+
assert response.json()["error"] == "invalid_hf_port"
|
| 402 |
+
|
| 403 |
+
|
| 404 |
+
def test_status_reports_direct_hf_ws_url_as_ready(
|
| 405 |
+
tmp_path: Path,
|
| 406 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 407 |
+
) -> None:
|
| 408 |
+
"""Settings API should treat a direct Hugging Face websocket as a valid configuration."""
|
| 409 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 410 |
+
monkeypatch.setattr(config, "HF_REALTIME_CONNECTION_MODE", "local")
|
| 411 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", None)
|
| 412 |
+
monkeypatch.setattr(config, "HF_REALTIME_WS_URL", "ws://127.0.0.1:8765/v1/realtime")
|
| 413 |
+
|
| 414 |
+
app = FastAPI()
|
| 415 |
+
robot = SimpleNamespace(media=SimpleNamespace(audio=None, backend=None))
|
| 416 |
+
stream = LocalStream(MagicMock(), robot, settings_app=app, instance_path=str(tmp_path))
|
| 417 |
+
stream._init_settings_ui_if_needed()
|
| 418 |
+
|
| 419 |
+
client = TestClient(app)
|
| 420 |
+
response = client.get("/status")
|
| 421 |
+
|
| 422 |
+
assert response.status_code == 200
|
| 423 |
+
data = response.json()
|
| 424 |
+
assert data["backend_provider"] == "huggingface"
|
| 425 |
+
assert data["has_hf_session_url"] is False
|
| 426 |
+
assert data["has_hf_ws_url"] is True
|
| 427 |
+
assert data["has_hf_connection"] is True
|
| 428 |
+
assert data["hf_connection_mode"] == "local"
|
| 429 |
+
assert data["can_proceed_with_hf"] is True
|
| 430 |
+
|
| 431 |
+
|
| 432 |
+
def test_headless_personality_routes_return_gemini_voices_when_backend_selected(
|
| 433 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 434 |
+
) -> None:
|
| 435 |
"""Headless personality UI should expose Gemini voices when Gemini is selected."""
|
| 436 |
monkeypatch.setattr(config, "BACKEND_PROVIDER", "gemini")
|
| 437 |
monkeypatch.setattr(config, "MODEL_NAME", "gemini-3.1-flash-live-preview")
|
|
|
|
| 447 |
assert response.json() == GEMINI_AVAILABLE_VOICES
|
| 448 |
|
| 449 |
|
| 450 |
+
def test_headless_personality_routes_load_builtin_default_tools() -> None:
|
| 451 |
+
"""Headless personality UI should expose built-in default tools on initial load."""
|
| 452 |
+
app = FastAPI()
|
| 453 |
+
handler = MagicMock()
|
| 454 |
+
mount_personality_routes(app, handler, lambda: None)
|
| 455 |
+
|
| 456 |
+
client = TestClient(app)
|
| 457 |
+
response = client.get("/personalities/load", params={"name": "(built-in default)"})
|
| 458 |
+
|
| 459 |
+
assert response.status_code == 200
|
| 460 |
+
data = response.json()
|
| 461 |
+
assert data["tools_text"]
|
| 462 |
+
assert "dance" in data["enabled_tools"]
|
| 463 |
+
assert "camera" in data["enabled_tools"]
|
| 464 |
+
|
| 465 |
+
|
| 466 |
def test_headless_personality_routes_apply_voice_accepts_query_param() -> None:
|
| 467 |
"""Headless personality UI should apply a voice change from a POST query param."""
|
| 468 |
app = FastAPI()
|
|
|
|
| 494 |
loop.call_soon_threadsafe(loop.stop)
|
| 495 |
thread.join(timeout=1.0)
|
| 496 |
loop.close()
|
| 497 |
+
|
| 498 |
+
|
| 499 |
+
def test_headless_personality_routes_persist_startup_with_voice_override() -> None:
|
| 500 |
+
"""Saving a startup personality should persist the active manual voice override."""
|
| 501 |
+
app = FastAPI()
|
| 502 |
+
handler = MagicMock()
|
| 503 |
+
handler.apply_personality = AsyncMock(return_value="Applied personality and restarted realtime session.")
|
| 504 |
+
handler.get_current_voice = MagicMock(return_value="shimmer")
|
| 505 |
+
persist_personality = MagicMock()
|
| 506 |
+
|
| 507 |
+
loop = asyncio.new_event_loop()
|
| 508 |
+
started = threading.Event()
|
| 509 |
+
|
| 510 |
+
def _run_loop() -> None:
|
| 511 |
+
asyncio.set_event_loop(loop)
|
| 512 |
+
started.set()
|
| 513 |
+
loop.run_forever()
|
| 514 |
+
|
| 515 |
+
thread = threading.Thread(target=_run_loop, daemon=True)
|
| 516 |
+
thread.start()
|
| 517 |
+
started.wait(timeout=1.0)
|
| 518 |
+
|
| 519 |
+
try:
|
| 520 |
+
mount_personality_routes(app, handler, lambda: loop, persist_personality=persist_personality)
|
| 521 |
+
|
| 522 |
+
client = TestClient(app)
|
| 523 |
+
response = client.post("/personalities/apply?name=sorry_bro&persist=1")
|
| 524 |
+
|
| 525 |
+
assert response.status_code == 200
|
| 526 |
+
assert response.json()["ok"] is True
|
| 527 |
+
handler.apply_personality.assert_awaited_once_with("sorry_bro")
|
| 528 |
+
persist_personality.assert_called_once_with("sorry_bro", "shimmer")
|
| 529 |
+
finally:
|
| 530 |
+
loop.call_soon_threadsafe(loop.stop)
|
| 531 |
+
thread.join(timeout=1.0)
|
| 532 |
+
loop.close()
|
| 533 |
+
|
| 534 |
+
|
| 535 |
+
def test_local_stream_persist_personality_stores_voice_override(tmp_path) -> None:
|
| 536 |
+
"""Persisting startup settings should write both profile and voice override."""
|
| 537 |
+
stream = LocalStream(MagicMock(), MagicMock(), instance_path=str(tmp_path))
|
| 538 |
+
|
| 539 |
+
stream._persist_personality("sorry_bro", "shimmer")
|
| 540 |
+
|
| 541 |
+
settings_path = tmp_path / "startup_settings.json"
|
| 542 |
+
assert settings_path.exists()
|
| 543 |
+
assert settings_path.read_text(encoding="utf-8") == '{\n "profile": "sorry_bro",\n "voice": "shimmer"\n}\n'
|
| 544 |
+
assert stream._read_persisted_personality() == "sorry_bro"
|
| 545 |
+
|
| 546 |
+
|
| 547 |
+
def test_local_stream_persist_personality_clears_legacy_startup_env_overrides(tmp_path, monkeypatch) -> None:
|
| 548 |
+
"""Saving startup settings should remove legacy `.env` profile and voice overrides."""
|
| 549 |
+
env_path = tmp_path / ".env"
|
| 550 |
+
env_path.write_text(
|
| 551 |
+
"OPENAI_API_KEY=test-key\n"
|
| 552 |
+
"REACHY_MINI_CUSTOM_PROFILE=mad_scientist_assistant\n"
|
| 553 |
+
"REACHY_MINI_VOICE_OVERRIDE=shimmer\n",
|
| 554 |
+
encoding="utf-8",
|
| 555 |
+
)
|
| 556 |
+
stream = LocalStream(MagicMock(), MagicMock(), instance_path=str(tmp_path))
|
| 557 |
+
|
| 558 |
+
stream._persist_personality(None, "Aiden")
|
| 559 |
+
|
| 560 |
+
env_text = env_path.read_text(encoding="utf-8")
|
| 561 |
+
assert "OPENAI_API_KEY=test-key" in env_text
|
| 562 |
+
assert "REACHY_MINI_CUSTOM_PROFILE=" not in env_text
|
| 563 |
+
assert "REACHY_MINI_VOICE_OVERRIDE=" not in env_text
|
| 564 |
+
|
| 565 |
+
applied_profiles: list[str | None] = []
|
| 566 |
+
monkeypatch.delenv("REACHY_MINI_CUSTOM_PROFILE", raising=False)
|
| 567 |
+
monkeypatch.setattr(
|
| 568 |
+
"reachy_mini_conversation_app.config.set_custom_profile",
|
| 569 |
+
lambda profile: applied_profiles.append(profile),
|
| 570 |
+
)
|
| 571 |
+
|
| 572 |
+
settings = load_startup_settings_into_runtime(tmp_path)
|
| 573 |
+
|
| 574 |
+
assert settings == StartupSettings(voice="Aiden")
|
| 575 |
+
assert applied_profiles == [None]
|
| 576 |
+
|
| 577 |
+
|
| 578 |
+
def test_local_stream_launch_waits_for_manual_openai_key_without_download(
|
| 579 |
+
tmp_path: Path,
|
| 580 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 581 |
+
) -> None:
|
| 582 |
+
"""OpenAI startup should wait for settings input instead of claiming a bundled key."""
|
| 583 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 584 |
+
monkeypatch.setattr(config, "OPENAI_API_KEY", None)
|
| 585 |
+
monkeypatch.setenv("BACKEND_PROVIDER", "openai")
|
| 586 |
+
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
| 587 |
+
|
| 588 |
+
fake_client_ctor = MagicMock(side_effect=AssertionError("launch() should not try to download an OpenAI key"))
|
| 589 |
+
monkeypatch.setitem(sys.modules, "gradio_client", SimpleNamespace(Client=fake_client_ctor))
|
| 590 |
+
|
| 591 |
+
media = SimpleNamespace(
|
| 592 |
+
start_recording=MagicMock(),
|
| 593 |
+
start_playing=MagicMock(),
|
| 594 |
+
)
|
| 595 |
+
robot = SimpleNamespace(media=media)
|
| 596 |
+
stream = LocalStream(MagicMock(), robot, settings_app=FastAPI(), instance_path=str(tmp_path))
|
| 597 |
+
stream._active_backend_name = "openai"
|
| 598 |
+
|
| 599 |
+
init_settings_ui = MagicMock()
|
| 600 |
+
monkeypatch.setattr(stream, "_init_settings_ui_if_needed", init_settings_ui)
|
| 601 |
+
monkeypatch.setattr(stream, "_has_required_key", MagicMock(side_effect=[False, False]))
|
| 602 |
+
monkeypatch.setattr("reachy_mini_conversation_app.console.time.sleep", MagicMock(side_effect=KeyboardInterrupt))
|
| 603 |
+
|
| 604 |
+
stream.launch()
|
| 605 |
+
|
| 606 |
+
fake_client_ctor.assert_not_called()
|
| 607 |
+
init_settings_ui.assert_called_once()
|
| 608 |
+
media.start_recording.assert_not_called()
|
| 609 |
+
media.start_playing.assert_not_called()
|
tests/test_gemini_live.py
CHANGED
|
@@ -3,6 +3,7 @@
|
|
| 3 |
import base64
|
| 4 |
import asyncio
|
| 5 |
from types import SimpleNamespace
|
|
|
|
| 6 |
from unittest.mock import AsyncMock, MagicMock, call
|
| 7 |
|
| 8 |
import numpy as np
|
|
@@ -10,13 +11,14 @@ import pytest
|
|
| 10 |
from fastrtc import AdditionalOutputs
|
| 11 |
|
| 12 |
import reachy_mini_conversation_app.gemini_live as gemini_mod
|
|
|
|
| 13 |
from reachy_mini_conversation_app.gemini_live import GeminiLiveHandler
|
| 14 |
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 15 |
from reachy_mini_conversation_app.tools.tool_constants import ToolState
|
| 16 |
from reachy_mini_conversation_app.tools.background_tool_manager import ToolNotification
|
| 17 |
|
| 18 |
|
| 19 |
-
def _server_content(**kwargs):
|
| 20 |
defaults = {
|
| 21 |
"model_turn": None,
|
| 22 |
"turn_complete": None,
|
|
@@ -33,11 +35,11 @@ def _server_content(**kwargs):
|
|
| 33 |
return SimpleNamespace(**defaults)
|
| 34 |
|
| 35 |
|
| 36 |
-
def _response(server_content=None, tool_call=None):
|
| 37 |
return SimpleNamespace(server_content=server_content, tool_call=tool_call)
|
| 38 |
|
| 39 |
|
| 40 |
-
async def _wait_for(predicate, timeout: float = 1.0) -> None:
|
| 41 |
deadline = asyncio.get_running_loop().time() + timeout
|
| 42 |
while asyncio.get_running_loop().time() < deadline:
|
| 43 |
if predicate():
|
|
@@ -47,24 +49,24 @@ async def _wait_for(predicate, timeout: float = 1.0) -> None:
|
|
| 47 |
|
| 48 |
|
| 49 |
class _FakeSession:
|
| 50 |
-
def __init__(self, batches, stop_event: asyncio.Event):
|
| 51 |
self._batches = list(batches)
|
| 52 |
self._stop_event = stop_event
|
| 53 |
-
self.realtime_inputs = []
|
| 54 |
-
self.tool_responses = []
|
| 55 |
|
| 56 |
async def close(self) -> None:
|
| 57 |
self._stop_event.set()
|
| 58 |
|
| 59 |
-
async def send_realtime_input(self, **kwargs) -> None:
|
| 60 |
self.realtime_inputs.append(kwargs)
|
| 61 |
return None
|
| 62 |
|
| 63 |
-
async def send_tool_response(self, **kwargs) -> None:
|
| 64 |
self.tool_responses.append(kwargs)
|
| 65 |
return None
|
| 66 |
|
| 67 |
-
async def receive(self):
|
| 68 |
if self._batches:
|
| 69 |
for response in self._batches.pop(0):
|
| 70 |
yield response
|
|
@@ -82,21 +84,23 @@ class _FakeConnectContext:
|
|
| 82 |
async def __aenter__(self) -> _FakeSession:
|
| 83 |
return self._session
|
| 84 |
|
| 85 |
-
async def __aexit__(self, *_args) -> bool:
|
| 86 |
return False
|
| 87 |
|
| 88 |
|
| 89 |
class _FakeLiveClient:
|
| 90 |
-
def __init__(self, session: _FakeSession):
|
| 91 |
self.aio = SimpleNamespace(live=SimpleNamespace(connect=lambda **_kwargs: _FakeConnectContext(session)))
|
| 92 |
|
| 93 |
|
| 94 |
@pytest.mark.asyncio
|
| 95 |
-
async def test_gemini_turn_buffers_transcripts_and_schedules_motion_reset(
|
|
|
|
|
|
|
| 96 |
"""Gemini turns should emit one transcript per role and let the wobbler reset after speech."""
|
| 97 |
monkeypatch.setattr(gemini_mod, "get_session_instructions", lambda: "test")
|
| 98 |
monkeypatch.setattr(gemini_mod, "get_session_voice", lambda: "Kore")
|
| 99 |
-
monkeypatch.setattr(gemini_mod, "
|
| 100 |
|
| 101 |
movement_manager = MagicMock()
|
| 102 |
movement_manager.is_idle.return_value = False
|
|
@@ -108,8 +112,8 @@ async def test_gemini_turn_buffers_transcripts_and_schedules_motion_reset(monkey
|
|
| 108 |
head_wobbler=head_wobbler,
|
| 109 |
)
|
| 110 |
handler = GeminiLiveHandler(deps)
|
| 111 |
-
|
| 112 |
-
|
| 113 |
|
| 114 |
audio_bytes = b"\x00\x00\x10\x00" * 256
|
| 115 |
session = _FakeSession(
|
|
@@ -235,3 +239,80 @@ async def test_gemini_camera_tool_sends_snapshot_and_returns_json_result() -> No
|
|
| 235 |
"metadata": {"title": "🛠️ Used tool camera", "status": "done"},
|
| 236 |
}
|
| 237 |
]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
import base64
|
| 4 |
import asyncio
|
| 5 |
from types import SimpleNamespace
|
| 6 |
+
from typing import Any, Callable, AsyncIterator
|
| 7 |
from unittest.mock import AsyncMock, MagicMock, call
|
| 8 |
|
| 9 |
import numpy as np
|
|
|
|
| 11 |
from fastrtc import AdditionalOutputs
|
| 12 |
|
| 13 |
import reachy_mini_conversation_app.gemini_live as gemini_mod
|
| 14 |
+
import reachy_mini_conversation_app.tools.core_tools as ct_mod
|
| 15 |
from reachy_mini_conversation_app.gemini_live import GeminiLiveHandler
|
| 16 |
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 17 |
from reachy_mini_conversation_app.tools.tool_constants import ToolState
|
| 18 |
from reachy_mini_conversation_app.tools.background_tool_manager import ToolNotification
|
| 19 |
|
| 20 |
|
| 21 |
+
def _server_content(**kwargs: Any) -> SimpleNamespace:
|
| 22 |
defaults = {
|
| 23 |
"model_turn": None,
|
| 24 |
"turn_complete": None,
|
|
|
|
| 35 |
return SimpleNamespace(**defaults)
|
| 36 |
|
| 37 |
|
| 38 |
+
def _response(server_content: Any = None, tool_call: Any = None) -> SimpleNamespace:
|
| 39 |
return SimpleNamespace(server_content=server_content, tool_call=tool_call)
|
| 40 |
|
| 41 |
|
| 42 |
+
async def _wait_for(predicate: Callable[[], bool], timeout: float = 1.0) -> None:
|
| 43 |
deadline = asyncio.get_running_loop().time() + timeout
|
| 44 |
while asyncio.get_running_loop().time() < deadline:
|
| 45 |
if predicate():
|
|
|
|
| 49 |
|
| 50 |
|
| 51 |
class _FakeSession:
|
| 52 |
+
def __init__(self, batches: list[list[SimpleNamespace]], stop_event: asyncio.Event) -> None:
|
| 53 |
self._batches = list(batches)
|
| 54 |
self._stop_event = stop_event
|
| 55 |
+
self.realtime_inputs: list[dict[str, Any]] = []
|
| 56 |
+
self.tool_responses: list[dict[str, Any]] = []
|
| 57 |
|
| 58 |
async def close(self) -> None:
|
| 59 |
self._stop_event.set()
|
| 60 |
|
| 61 |
+
async def send_realtime_input(self, **kwargs: Any) -> None:
|
| 62 |
self.realtime_inputs.append(kwargs)
|
| 63 |
return None
|
| 64 |
|
| 65 |
+
async def send_tool_response(self, **kwargs: Any) -> None:
|
| 66 |
self.tool_responses.append(kwargs)
|
| 67 |
return None
|
| 68 |
|
| 69 |
+
async def receive(self) -> AsyncIterator[SimpleNamespace]:
|
| 70 |
if self._batches:
|
| 71 |
for response in self._batches.pop(0):
|
| 72 |
yield response
|
|
|
|
| 84 |
async def __aenter__(self) -> _FakeSession:
|
| 85 |
return self._session
|
| 86 |
|
| 87 |
+
async def __aexit__(self, *_args: object) -> bool:
|
| 88 |
return False
|
| 89 |
|
| 90 |
|
| 91 |
class _FakeLiveClient:
|
| 92 |
+
def __init__(self, session: _FakeSession) -> None:
|
| 93 |
self.aio = SimpleNamespace(live=SimpleNamespace(connect=lambda **_kwargs: _FakeConnectContext(session)))
|
| 94 |
|
| 95 |
|
| 96 |
@pytest.mark.asyncio
|
| 97 |
+
async def test_gemini_turn_buffers_transcripts_and_schedules_motion_reset(
|
| 98 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 99 |
+
) -> None:
|
| 100 |
"""Gemini turns should emit one transcript per role and let the wobbler reset after speech."""
|
| 101 |
monkeypatch.setattr(gemini_mod, "get_session_instructions", lambda: "test")
|
| 102 |
monkeypatch.setattr(gemini_mod, "get_session_voice", lambda: "Kore")
|
| 103 |
+
monkeypatch.setattr(gemini_mod, "get_active_tool_specs", lambda _: [])
|
| 104 |
|
| 105 |
movement_manager = MagicMock()
|
| 106 |
movement_manager.is_idle.return_value = False
|
|
|
|
| 112 |
head_wobbler=head_wobbler,
|
| 113 |
)
|
| 114 |
handler = GeminiLiveHandler(deps)
|
| 115 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_up", MagicMock())
|
| 116 |
+
monkeypatch.setattr(type(handler.tool_manager), "shutdown", AsyncMock())
|
| 117 |
|
| 118 |
audio_bytes = b"\x00\x00\x10\x00" * 256
|
| 119 |
session = _FakeSession(
|
|
|
|
| 239 |
"metadata": {"title": "🛠️ Used tool camera", "status": "done"},
|
| 240 |
}
|
| 241 |
]
|
| 242 |
+
|
| 243 |
+
|
| 244 |
+
@pytest.mark.asyncio
|
| 245 |
+
async def test_apply_personality_preserves_manual_voice_override(monkeypatch) -> None:
|
| 246 |
+
"""Applying a profile should keep a manually selected Gemini voice active."""
|
| 247 |
+
monkeypatch.setattr(gemini_mod, "get_session_instructions", lambda: "test")
|
| 248 |
+
monkeypatch.setattr(gemini_mod, "get_session_voice", lambda: "Kore")
|
| 249 |
+
monkeypatch.setattr("reachy_mini_conversation_app.config.set_custom_profile", lambda _profile: None)
|
| 250 |
+
|
| 251 |
+
handler = GeminiLiveHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 252 |
+
handler.session = object()
|
| 253 |
+
handler._voice_override = "Orus"
|
| 254 |
+
restart = AsyncMock()
|
| 255 |
+
monkeypatch.setattr(handler, "_restart_session", restart)
|
| 256 |
+
|
| 257 |
+
status = await handler.apply_personality("example")
|
| 258 |
+
|
| 259 |
+
assert status == "Applied personality and restarted Gemini session."
|
| 260 |
+
assert handler.get_current_voice() == "Orus"
|
| 261 |
+
restart.assert_awaited_once()
|
| 262 |
+
|
| 263 |
+
|
| 264 |
+
def test_handler_uses_startup_voice_at_startup() -> None:
|
| 265 |
+
"""Gemini handler startup should restore a persisted startup voice."""
|
| 266 |
+
handler = GeminiLiveHandler(
|
| 267 |
+
ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()),
|
| 268 |
+
startup_voice="Orus",
|
| 269 |
+
)
|
| 270 |
+
|
| 271 |
+
assert handler.get_current_voice() == "Orus"
|
| 272 |
+
|
| 273 |
+
|
| 274 |
+
def test_copy_preserves_current_voice_override() -> None:
|
| 275 |
+
"""Copied Gemini handlers should keep the current voice override."""
|
| 276 |
+
handler = GeminiLiveHandler(
|
| 277 |
+
ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()),
|
| 278 |
+
startup_voice="Orus",
|
| 279 |
+
)
|
| 280 |
+
handler._voice_override = "Zephyr"
|
| 281 |
+
|
| 282 |
+
copied_handler = handler.copy()
|
| 283 |
+
|
| 284 |
+
assert copied_handler.get_current_voice() == "Zephyr"
|
| 285 |
+
|
| 286 |
+
|
| 287 |
+
def test_gemini_excludes_head_tracking_when_no_head_tracker(monkeypatch) -> None:
|
| 288 |
+
"""head_tracking tool must not appear in Gemini session config when head_tracker is not active."""
|
| 289 |
+
monkeypatch.setattr(gemini_mod, "get_session_instructions", lambda: "test")
|
| 290 |
+
monkeypatch.setattr(gemini_mod, "get_session_voice", lambda: "Kore")
|
| 291 |
+
|
| 292 |
+
# mock ALL_TOOL_SPECS to include at least head_tracking and one other tool, to verify that only head_tracking is excluded, not all tools
|
| 293 |
+
monkeypatch.setattr(
|
| 294 |
+
ct_mod,
|
| 295 |
+
"ALL_TOOL_SPECS",
|
| 296 |
+
[
|
| 297 |
+
{"type": "function", "name": "head_tracking", "description": "head_tracking", "parameters": {}},
|
| 298 |
+
{"type": "function", "name": "fake_tool", "description": "fake_tool", "parameters": {}},
|
| 299 |
+
],
|
| 300 |
+
)
|
| 301 |
+
|
| 302 |
+
# case 1: no camera at all, --no-camera flag passed
|
| 303 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock(), camera_worker=None)
|
| 304 |
+
handler = GeminiLiveHandler(deps)
|
| 305 |
+
live_config = handler._build_live_config()
|
| 306 |
+
tool_names = [fd.name for fd in live_config.tools[0].function_declarations] if live_config.tools else []
|
| 307 |
+
assert "head_tracking" not in tool_names, "case 1 failed: camera_worker=None"
|
| 308 |
+
assert "fake_tool" in tool_names, "case 1 failed: a non-head-tracking tool was unexpectedly excluded"
|
| 309 |
+
|
| 310 |
+
# case 2: camera is running but --head-tracker flag was not passed
|
| 311 |
+
camera_worker = MagicMock()
|
| 312 |
+
camera_worker.head_tracker = None
|
| 313 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock(), camera_worker=camera_worker)
|
| 314 |
+
handler = GeminiLiveHandler(deps)
|
| 315 |
+
live_config = handler._build_live_config()
|
| 316 |
+
tool_names = [fd.name for fd in live_config.tools[0].function_declarations] if live_config.tools else []
|
| 317 |
+
assert "head_tracking" not in tool_names, "case 2 failed: camera_worker.head_tracker=None"
|
| 318 |
+
assert "fake_tool" in tool_names, "case 2 failed: a non-head-tracking tool was unexpectedly excluded"
|
tests/test_huggingface_realtime.py
ADDED
|
@@ -0,0 +1,628 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import asyncio
|
| 2 |
+
from typing import Any
|
| 3 |
+
from unittest.mock import AsyncMock, MagicMock
|
| 4 |
+
|
| 5 |
+
import pytest
|
| 6 |
+
|
| 7 |
+
import reachy_mini_conversation_app.base_realtime as base_rt_mod
|
| 8 |
+
import reachy_mini_conversation_app.huggingface_realtime as hf_mod
|
| 9 |
+
from reachy_mini_conversation_app.config import HF_BACKEND, config, get_default_voice_for_backend
|
| 10 |
+
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 11 |
+
from reachy_mini_conversation_app.huggingface_realtime import HuggingFaceRealtimeHandler
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
HF_DEFAULT_VOICE = get_default_voice_for_backend(HF_BACKEND)
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def _make_usage(
|
| 18 |
+
audio_in: int | None = 100,
|
| 19 |
+
text_in: int | None = 200,
|
| 20 |
+
image_in: int | None = 300,
|
| 21 |
+
audio_out: int | None = 400,
|
| 22 |
+
text_out: int | None = 500,
|
| 23 |
+
has_input: bool = True,
|
| 24 |
+
has_output: bool = True,
|
| 25 |
+
) -> MagicMock:
|
| 26 |
+
"""Build a fake usage object matching the OpenAI-compatible response.usage shape."""
|
| 27 |
+
usage = MagicMock()
|
| 28 |
+
if has_input:
|
| 29 |
+
inp = MagicMock()
|
| 30 |
+
inp.audio_tokens = audio_in
|
| 31 |
+
inp.text_tokens = text_in
|
| 32 |
+
inp.image_tokens = image_in
|
| 33 |
+
usage.input_token_details = inp
|
| 34 |
+
else:
|
| 35 |
+
usage.input_token_details = None
|
| 36 |
+
if has_output:
|
| 37 |
+
out = MagicMock()
|
| 38 |
+
out.audio_tokens = audio_out
|
| 39 |
+
out.text_tokens = text_out
|
| 40 |
+
usage.output_token_details = out
|
| 41 |
+
else:
|
| 42 |
+
usage.output_token_details = None
|
| 43 |
+
return usage
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
@pytest.mark.asyncio
|
| 47 |
+
async def test_partial_transcription_uses_latest_snapshot(monkeypatch: Any) -> None:
|
| 48 |
+
"""Partial transcription snapshots should replace older snapshots for the same item."""
|
| 49 |
+
monkeypatch.setattr(hf_mod, "get_session_instructions", lambda: "test")
|
| 50 |
+
monkeypatch.setattr(hf_mod, "get_session_voice", lambda default=HF_DEFAULT_VOICE: "Aiden")
|
| 51 |
+
monkeypatch.setattr(hf_mod, "get_active_tool_specs", lambda _: [])
|
| 52 |
+
|
| 53 |
+
class FakeEvent:
|
| 54 |
+
def __init__(self, etype: str, **kwargs: Any) -> None:
|
| 55 |
+
self.type = etype
|
| 56 |
+
for key, value in kwargs.items():
|
| 57 |
+
setattr(self, key, value)
|
| 58 |
+
|
| 59 |
+
class FakeSession:
|
| 60 |
+
async def update(self, **_kw: Any) -> None:
|
| 61 |
+
pass
|
| 62 |
+
|
| 63 |
+
class FakeInputAudioBuffer:
|
| 64 |
+
async def append(self, **_kw: Any) -> None:
|
| 65 |
+
pass
|
| 66 |
+
|
| 67 |
+
class FakeItem:
|
| 68 |
+
async def create(self, **_kw: Any) -> None:
|
| 69 |
+
pass
|
| 70 |
+
|
| 71 |
+
class FakeConversation:
|
| 72 |
+
item = FakeItem()
|
| 73 |
+
|
| 74 |
+
class FakeResponse:
|
| 75 |
+
async def create(self, **_kw: Any) -> None:
|
| 76 |
+
pass
|
| 77 |
+
|
| 78 |
+
async def cancel(self, **_kw: Any) -> None:
|
| 79 |
+
pass
|
| 80 |
+
|
| 81 |
+
class FakeConn:
|
| 82 |
+
session = FakeSession()
|
| 83 |
+
input_audio_buffer = FakeInputAudioBuffer()
|
| 84 |
+
conversation = FakeConversation()
|
| 85 |
+
response = FakeResponse()
|
| 86 |
+
|
| 87 |
+
def __init__(self) -> None:
|
| 88 |
+
self._events = iter(
|
| 89 |
+
[
|
| 90 |
+
FakeEvent("conversation.item.input_audio_transcription.delta", item_id="item-1", delta="Hey"),
|
| 91 |
+
FakeEvent(
|
| 92 |
+
"conversation.item.input_audio_transcription.delta",
|
| 93 |
+
item_id="item-1",
|
| 94 |
+
delta="Hey, how are you?",
|
| 95 |
+
),
|
| 96 |
+
]
|
| 97 |
+
)
|
| 98 |
+
|
| 99 |
+
async def __aenter__(self) -> "FakeConn":
|
| 100 |
+
return self
|
| 101 |
+
|
| 102 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 103 |
+
return False
|
| 104 |
+
|
| 105 |
+
async def close(self) -> None:
|
| 106 |
+
pass
|
| 107 |
+
|
| 108 |
+
def __aiter__(self) -> "FakeConn":
|
| 109 |
+
return self
|
| 110 |
+
|
| 111 |
+
async def __anext__(self) -> FakeEvent:
|
| 112 |
+
try:
|
| 113 |
+
return next(self._events)
|
| 114 |
+
except StopIteration:
|
| 115 |
+
raise StopAsyncIteration
|
| 116 |
+
|
| 117 |
+
class FakeRealtime:
|
| 118 |
+
def connect(self, **_kw: Any) -> FakeConn:
|
| 119 |
+
return FakeConn()
|
| 120 |
+
|
| 121 |
+
class FakeClient:
|
| 122 |
+
def __init__(self) -> None:
|
| 123 |
+
self.realtime = FakeRealtime()
|
| 124 |
+
|
| 125 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 126 |
+
handler = HuggingFaceRealtimeHandler(deps)
|
| 127 |
+
fake_client = FakeClient()
|
| 128 |
+
handler.client = fake_client
|
| 129 |
+
|
| 130 |
+
start_up = MagicMock()
|
| 131 |
+
shutdown = AsyncMock()
|
| 132 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_up", start_up)
|
| 133 |
+
monkeypatch.setattr(type(handler.tool_manager), "shutdown", shutdown)
|
| 134 |
+
|
| 135 |
+
await handler._run_realtime_session()
|
| 136 |
+
|
| 137 |
+
assert handler.input_transcript_chunks_by_item.item_id == "item-1"
|
| 138 |
+
assert handler.input_transcript_chunks_by_item.deltas == ["Hey, how are you?"]
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
@pytest.mark.asyncio
|
| 142 |
+
async def test_output_audio_delta_passes_output_sample_rate_to_head_wobbler(monkeypatch: Any) -> None:
|
| 143 |
+
"""Assistant audio deltas should propagate the realtime output sample rate to the head wobbler."""
|
| 144 |
+
monkeypatch.setattr(hf_mod, "get_session_instructions", lambda: "test")
|
| 145 |
+
monkeypatch.setattr(hf_mod, "get_session_voice", lambda default=HF_DEFAULT_VOICE: "Aiden")
|
| 146 |
+
monkeypatch.setattr(hf_mod, "get_active_tool_specs", lambda _: [])
|
| 147 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 148 |
+
|
| 149 |
+
audio_delta = "AAABAAIAAwA="
|
| 150 |
+
|
| 151 |
+
class FakeEvent:
|
| 152 |
+
def __init__(self, etype: str, **kwargs: Any) -> None:
|
| 153 |
+
self.type = etype
|
| 154 |
+
for key, value in kwargs.items():
|
| 155 |
+
setattr(self, key, value)
|
| 156 |
+
|
| 157 |
+
class FakeSession:
|
| 158 |
+
async def update(self, **_kw: Any) -> None:
|
| 159 |
+
pass
|
| 160 |
+
|
| 161 |
+
class FakeInputAudioBuffer:
|
| 162 |
+
async def append(self, **_kw: Any) -> None:
|
| 163 |
+
pass
|
| 164 |
+
|
| 165 |
+
class FakeItem:
|
| 166 |
+
async def create(self, **_kw: Any) -> None:
|
| 167 |
+
pass
|
| 168 |
+
|
| 169 |
+
class FakeConversation:
|
| 170 |
+
item = FakeItem()
|
| 171 |
+
|
| 172 |
+
class FakeResponse:
|
| 173 |
+
async def create(self, **_kw: Any) -> None:
|
| 174 |
+
pass
|
| 175 |
+
|
| 176 |
+
async def cancel(self, **_kw: Any) -> None:
|
| 177 |
+
pass
|
| 178 |
+
|
| 179 |
+
class FakeConn:
|
| 180 |
+
session = FakeSession()
|
| 181 |
+
input_audio_buffer = FakeInputAudioBuffer()
|
| 182 |
+
conversation = FakeConversation()
|
| 183 |
+
response = FakeResponse()
|
| 184 |
+
|
| 185 |
+
def __init__(self) -> None:
|
| 186 |
+
self._events = iter(
|
| 187 |
+
[
|
| 188 |
+
FakeEvent("response.created"),
|
| 189 |
+
FakeEvent("response.output_audio.delta", delta=audio_delta),
|
| 190 |
+
FakeEvent("response.output_audio.done"),
|
| 191 |
+
]
|
| 192 |
+
)
|
| 193 |
+
|
| 194 |
+
async def __aenter__(self) -> "FakeConn":
|
| 195 |
+
return self
|
| 196 |
+
|
| 197 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 198 |
+
return False
|
| 199 |
+
|
| 200 |
+
async def close(self) -> None:
|
| 201 |
+
pass
|
| 202 |
+
|
| 203 |
+
def __aiter__(self) -> "FakeConn":
|
| 204 |
+
return self
|
| 205 |
+
|
| 206 |
+
async def __anext__(self) -> FakeEvent:
|
| 207 |
+
try:
|
| 208 |
+
return next(self._events)
|
| 209 |
+
except StopIteration:
|
| 210 |
+
raise StopAsyncIteration
|
| 211 |
+
|
| 212 |
+
class FakeRealtime:
|
| 213 |
+
def connect(self, **_kw: Any) -> FakeConn:
|
| 214 |
+
return FakeConn()
|
| 215 |
+
|
| 216 |
+
class FakeClient:
|
| 217 |
+
def __init__(self) -> None:
|
| 218 |
+
self.realtime = FakeRealtime()
|
| 219 |
+
|
| 220 |
+
head_wobbler = MagicMock()
|
| 221 |
+
deps = ToolDependencies(
|
| 222 |
+
reachy_mini=MagicMock(),
|
| 223 |
+
movement_manager=MagicMock(),
|
| 224 |
+
head_wobbler=head_wobbler,
|
| 225 |
+
)
|
| 226 |
+
handler = HuggingFaceRealtimeHandler(deps, gradio_mode=True)
|
| 227 |
+
handler.client = FakeClient()
|
| 228 |
+
|
| 229 |
+
start_up = MagicMock()
|
| 230 |
+
shutdown = AsyncMock()
|
| 231 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_up", start_up)
|
| 232 |
+
monkeypatch.setattr(type(handler.tool_manager), "shutdown", shutdown)
|
| 233 |
+
|
| 234 |
+
await handler._run_realtime_session()
|
| 235 |
+
|
| 236 |
+
head_wobbler.feed_pcm.assert_called_once()
|
| 237 |
+
assert head_wobbler.feed_pcm.call_args.args[1] == handler.output_sample_rate
|
| 238 |
+
head_wobbler.request_reset_after_current_audio.assert_called_once()
|
| 239 |
+
|
| 240 |
+
|
| 241 |
+
@pytest.mark.asyncio
|
| 242 |
+
async def test_emit_skips_idle_signal_while_response_active(monkeypatch: Any) -> None:
|
| 243 |
+
"""Idle tools should not trigger while a response is still active."""
|
| 244 |
+
movement_manager = MagicMock()
|
| 245 |
+
movement_manager.is_idle.return_value = True
|
| 246 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=movement_manager)
|
| 247 |
+
handler = HuggingFaceRealtimeHandler(deps)
|
| 248 |
+
handler.last_activity_time = asyncio.get_running_loop().time() - 60.0
|
| 249 |
+
handler._response_done_event.clear()
|
| 250 |
+
|
| 251 |
+
send_idle_signal = AsyncMock()
|
| 252 |
+
monkeypatch.setattr(handler, "send_idle_signal", send_idle_signal)
|
| 253 |
+
monkeypatch.setattr(base_rt_mod, "wait_for_item", AsyncMock(return_value=None))
|
| 254 |
+
|
| 255 |
+
result = await handler.emit()
|
| 256 |
+
|
| 257 |
+
assert result is None
|
| 258 |
+
send_idle_signal.assert_not_awaited()
|
| 259 |
+
|
| 260 |
+
|
| 261 |
+
def test_handler_uses_hf_startup_voice_at_startup(monkeypatch: Any) -> None:
|
| 262 |
+
"""Hugging Face startup should restore persisted HF voices."""
|
| 263 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 264 |
+
|
| 265 |
+
handler = HuggingFaceRealtimeHandler(
|
| 266 |
+
ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()),
|
| 267 |
+
startup_voice="Aiden",
|
| 268 |
+
)
|
| 269 |
+
|
| 270 |
+
assert handler.get_current_voice() == "Aiden"
|
| 271 |
+
|
| 272 |
+
|
| 273 |
+
@pytest.mark.asyncio
|
| 274 |
+
async def test_start_up_hf_gradio_does_not_wait_for_api_key(monkeypatch: Any) -> None:
|
| 275 |
+
"""Hugging Face backend should not wait for gradio key input."""
|
| 276 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 277 |
+
monkeypatch.setattr(config, "OPENAI_API_KEY", "sk-openai-secret")
|
| 278 |
+
|
| 279 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 280 |
+
handler = hf_mod.HuggingFaceRealtimeHandler(deps, gradio_mode=True)
|
| 281 |
+
|
| 282 |
+
build_client = AsyncMock(return_value=MagicMock())
|
| 283 |
+
run_realtime_session = AsyncMock(return_value=None)
|
| 284 |
+
wait_for_args = AsyncMock(side_effect=AssertionError("wait_for_args should not be called"))
|
| 285 |
+
|
| 286 |
+
monkeypatch.setattr(handler, "_build_realtime_client", build_client)
|
| 287 |
+
monkeypatch.setattr(handler, "_run_realtime_session", run_realtime_session)
|
| 288 |
+
monkeypatch.setattr(handler, "wait_for_args", wait_for_args)
|
| 289 |
+
|
| 290 |
+
await handler.start_up()
|
| 291 |
+
|
| 292 |
+
wait_for_args.assert_not_awaited()
|
| 293 |
+
build_client.assert_awaited_once_with()
|
| 294 |
+
run_realtime_session.assert_awaited_once()
|
| 295 |
+
|
| 296 |
+
|
| 297 |
+
@pytest.mark.asyncio
|
| 298 |
+
async def test_run_realtime_session_uses_default_voice_for_lb_allocated_sessions(monkeypatch: Any) -> None:
|
| 299 |
+
"""Use the backend default speaker when no profile voice is selected for the hf LB."""
|
| 300 |
+
monkeypatch.setattr(hf_mod, "get_session_instructions", lambda: "test")
|
| 301 |
+
monkeypatch.setattr(hf_mod, "get_session_voice", lambda default=HF_DEFAULT_VOICE: default)
|
| 302 |
+
monkeypatch.setattr(hf_mod, "get_active_tool_specs", lambda _: [])
|
| 303 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 304 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", "https://lb.example.test/session")
|
| 305 |
+
|
| 306 |
+
captured_update: dict[str, Any] = {}
|
| 307 |
+
|
| 308 |
+
class FakeSession:
|
| 309 |
+
async def update(self, **kwargs: Any) -> None:
|
| 310 |
+
captured_update.update(kwargs)
|
| 311 |
+
|
| 312 |
+
class FakeInputAudioBuffer:
|
| 313 |
+
async def append(self, **_kw: Any) -> None:
|
| 314 |
+
pass
|
| 315 |
+
|
| 316 |
+
class FakeItem:
|
| 317 |
+
async def create(self, **_kw: Any) -> None:
|
| 318 |
+
pass
|
| 319 |
+
|
| 320 |
+
class FakeConversation:
|
| 321 |
+
item = FakeItem()
|
| 322 |
+
|
| 323 |
+
class FakeResponse:
|
| 324 |
+
async def create(self, **_kw: Any) -> None:
|
| 325 |
+
pass
|
| 326 |
+
|
| 327 |
+
async def cancel(self, **_kw: Any) -> None:
|
| 328 |
+
pass
|
| 329 |
+
|
| 330 |
+
class FakeConn:
|
| 331 |
+
session = FakeSession()
|
| 332 |
+
input_audio_buffer = FakeInputAudioBuffer()
|
| 333 |
+
conversation = FakeConversation()
|
| 334 |
+
response = FakeResponse()
|
| 335 |
+
|
| 336 |
+
async def __aenter__(self) -> "FakeConn":
|
| 337 |
+
return self
|
| 338 |
+
|
| 339 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 340 |
+
return False
|
| 341 |
+
|
| 342 |
+
async def close(self) -> None:
|
| 343 |
+
pass
|
| 344 |
+
|
| 345 |
+
def __aiter__(self) -> "FakeConn":
|
| 346 |
+
return self
|
| 347 |
+
|
| 348 |
+
async def __anext__(self) -> Any:
|
| 349 |
+
raise StopAsyncIteration
|
| 350 |
+
|
| 351 |
+
class FakeRealtime:
|
| 352 |
+
def connect(self, **_kw: Any) -> FakeConn:
|
| 353 |
+
return FakeConn()
|
| 354 |
+
|
| 355 |
+
class FakeClient:
|
| 356 |
+
def __init__(self) -> None:
|
| 357 |
+
self.realtime = FakeRealtime()
|
| 358 |
+
|
| 359 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 360 |
+
handler = HuggingFaceRealtimeHandler(deps)
|
| 361 |
+
fake_client = FakeClient()
|
| 362 |
+
handler.client = fake_client
|
| 363 |
+
|
| 364 |
+
await handler._run_realtime_session()
|
| 365 |
+
|
| 366 |
+
session = captured_update["session"]
|
| 367 |
+
# HF at 16 kHz passes None so the backend uses its optimal default (16 kHz).
|
| 368 |
+
assert session["audio"]["input"]["format"]["rate"] is None
|
| 369 |
+
assert session["audio"]["output"]["format"]["rate"] is None
|
| 370 |
+
output = session["audio"]["output"]
|
| 371 |
+
assert output["voice"] == HF_DEFAULT_VOICE
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
@pytest.mark.asyncio
|
| 375 |
+
async def test_run_realtime_session_passes_allocated_session_query(monkeypatch: Any) -> None:
|
| 376 |
+
"""Hugging Face sessions must forward the allocated session token to the websocket connect call."""
|
| 377 |
+
monkeypatch.setattr(hf_mod, "get_session_instructions", lambda: "test")
|
| 378 |
+
monkeypatch.setattr(hf_mod, "get_session_voice", lambda default=HF_DEFAULT_VOICE: default)
|
| 379 |
+
monkeypatch.setattr(hf_mod, "get_active_tool_specs", lambda _: [])
|
| 380 |
+
|
| 381 |
+
captured_connect: dict[str, Any] = {}
|
| 382 |
+
|
| 383 |
+
class FakeSession:
|
| 384 |
+
async def update(self, **_kw: Any) -> None:
|
| 385 |
+
pass
|
| 386 |
+
|
| 387 |
+
class FakeInputAudioBuffer:
|
| 388 |
+
async def append(self, **_kw: Any) -> None:
|
| 389 |
+
pass
|
| 390 |
+
|
| 391 |
+
class FakeItem:
|
| 392 |
+
async def create(self, **_kw: Any) -> None:
|
| 393 |
+
pass
|
| 394 |
+
|
| 395 |
+
class FakeConversation:
|
| 396 |
+
item = FakeItem()
|
| 397 |
+
|
| 398 |
+
class FakeResponse:
|
| 399 |
+
async def create(self, **_kw: Any) -> None:
|
| 400 |
+
pass
|
| 401 |
+
|
| 402 |
+
async def cancel(self, **_kw: Any) -> None:
|
| 403 |
+
pass
|
| 404 |
+
|
| 405 |
+
class FakeConn:
|
| 406 |
+
session = FakeSession()
|
| 407 |
+
input_audio_buffer = FakeInputAudioBuffer()
|
| 408 |
+
conversation = FakeConversation()
|
| 409 |
+
response = FakeResponse()
|
| 410 |
+
|
| 411 |
+
async def __aenter__(self) -> "FakeConn":
|
| 412 |
+
return self
|
| 413 |
+
|
| 414 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 415 |
+
return False
|
| 416 |
+
|
| 417 |
+
async def close(self) -> None:
|
| 418 |
+
pass
|
| 419 |
+
|
| 420 |
+
def __aiter__(self) -> "FakeConn":
|
| 421 |
+
return self
|
| 422 |
+
|
| 423 |
+
async def __anext__(self) -> Any:
|
| 424 |
+
raise StopAsyncIteration
|
| 425 |
+
|
| 426 |
+
class FakeRealtime:
|
| 427 |
+
def connect(self, **kwargs: Any) -> FakeConn:
|
| 428 |
+
captured_connect.update(kwargs)
|
| 429 |
+
return FakeConn()
|
| 430 |
+
|
| 431 |
+
class FakeClient:
|
| 432 |
+
def __init__(self) -> None:
|
| 433 |
+
self.realtime = FakeRealtime()
|
| 434 |
+
|
| 435 |
+
handler = HuggingFaceRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 436 |
+
fake_client = FakeClient()
|
| 437 |
+
handler.client = fake_client
|
| 438 |
+
handler._realtime_connect_query = {"session_token": "abc123"}
|
| 439 |
+
|
| 440 |
+
await handler._run_realtime_session()
|
| 441 |
+
|
| 442 |
+
assert "model" not in captured_connect
|
| 443 |
+
assert captured_connect["extra_query"] == {"session_token": "abc123"}
|
| 444 |
+
|
| 445 |
+
|
| 446 |
+
@pytest.mark.asyncio
|
| 447 |
+
async def test_build_realtime_client_uses_direct_hf_ws_url(monkeypatch: Any) -> None:
|
| 448 |
+
"""Hugging Face direct websocket mode should bypass the session allocator."""
|
| 449 |
+
captured_client_kwargs: dict[str, Any] = {}
|
| 450 |
+
|
| 451 |
+
class FakeClient:
|
| 452 |
+
def __init__(self, **kwargs: Any) -> None:
|
| 453 |
+
captured_client_kwargs.update(kwargs)
|
| 454 |
+
|
| 455 |
+
def _unexpected_async_client(*_args: Any, **_kwargs: Any) -> Any:
|
| 456 |
+
raise AssertionError("session allocator should not be called in direct websocket mode")
|
| 457 |
+
|
| 458 |
+
monkeypatch.setattr(hf_mod, "AsyncOpenAI", FakeClient)
|
| 459 |
+
monkeypatch.setattr(hf_mod.httpx, "AsyncClient", _unexpected_async_client)
|
| 460 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 461 |
+
monkeypatch.setattr(config, "HF_REALTIME_CONNECTION_MODE", "local")
|
| 462 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", "https://lb.example.test/session")
|
| 463 |
+
monkeypatch.setattr(config, "OPENAI_API_KEY", "sk-openai-secret")
|
| 464 |
+
monkeypatch.setattr(config, "HF_TOKEN", None)
|
| 465 |
+
monkeypatch.setattr(
|
| 466 |
+
config,
|
| 467 |
+
"HF_REALTIME_WS_URL",
|
| 468 |
+
"ws://127.0.0.1:8765/v1/realtime?session_token=abc123&model=ignored-by-sdk",
|
| 469 |
+
)
|
| 470 |
+
|
| 471 |
+
handler = HuggingFaceRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 472 |
+
|
| 473 |
+
client = await handler._build_realtime_client()
|
| 474 |
+
|
| 475 |
+
assert client is not None
|
| 476 |
+
assert captured_client_kwargs["api_key"] == "DUMMY"
|
| 477 |
+
assert captured_client_kwargs["base_url"] == "http://127.0.0.1:8765/v1"
|
| 478 |
+
assert captured_client_kwargs["websocket_base_url"] == "ws://127.0.0.1:8765/v1"
|
| 479 |
+
assert handler._realtime_connect_query == {"session_token": "abc123"}
|
| 480 |
+
|
| 481 |
+
|
| 482 |
+
@pytest.mark.asyncio
|
| 483 |
+
async def test_build_realtime_client_uses_deployed_mode_even_when_direct_hf_ws_url_is_saved(
|
| 484 |
+
monkeypatch: Any,
|
| 485 |
+
) -> None:
|
| 486 |
+
"""Explicit deployed mode should let .env recover from a stale local websocket URL."""
|
| 487 |
+
captured_client_kwargs: dict[str, Any] = {}
|
| 488 |
+
requested_session_urls: list[str] = []
|
| 489 |
+
requested_session_headers: list[dict[str, str] | None] = []
|
| 490 |
+
|
| 491 |
+
class FakeClient:
|
| 492 |
+
def __init__(self, **kwargs: Any) -> None:
|
| 493 |
+
captured_client_kwargs.update(kwargs)
|
| 494 |
+
|
| 495 |
+
class FakeResponse:
|
| 496 |
+
def raise_for_status(self) -> None:
|
| 497 |
+
pass
|
| 498 |
+
|
| 499 |
+
def json(self) -> dict[str, str]:
|
| 500 |
+
return {
|
| 501 |
+
"session_id": "session-123",
|
| 502 |
+
"connect_url": "wss://hf.example.test/v1/realtime?session_token=allocated",
|
| 503 |
+
}
|
| 504 |
+
|
| 505 |
+
class FakeAsyncClient:
|
| 506 |
+
def __init__(self, **_kwargs: Any) -> None:
|
| 507 |
+
pass
|
| 508 |
+
|
| 509 |
+
async def __aenter__(self) -> "FakeAsyncClient":
|
| 510 |
+
return self
|
| 511 |
+
|
| 512 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 513 |
+
return False
|
| 514 |
+
|
| 515 |
+
async def post(self, url: str, headers: dict[str, str] | None = None) -> FakeResponse:
|
| 516 |
+
requested_session_urls.append(url)
|
| 517 |
+
requested_session_headers.append(headers)
|
| 518 |
+
return FakeResponse()
|
| 519 |
+
|
| 520 |
+
monkeypatch.setattr(hf_mod, "AsyncOpenAI", FakeClient)
|
| 521 |
+
monkeypatch.setattr(hf_mod.httpx, "AsyncClient", FakeAsyncClient)
|
| 522 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 523 |
+
monkeypatch.setattr(config, "HF_REALTIME_CONNECTION_MODE", "deployed")
|
| 524 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", "https://lb.example.test/session")
|
| 525 |
+
monkeypatch.setattr(config, "HF_REALTIME_WS_URL", "ws://127.0.0.1:8765/v1/realtime")
|
| 526 |
+
monkeypatch.setattr(config, "OPENAI_API_KEY", "sk-openai-secret")
|
| 527 |
+
monkeypatch.setattr(config, "HF_TOKEN", "hf-secret")
|
| 528 |
+
|
| 529 |
+
handler = HuggingFaceRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 530 |
+
|
| 531 |
+
client = await handler._build_realtime_client()
|
| 532 |
+
|
| 533 |
+
assert client is not None
|
| 534 |
+
assert requested_session_urls == ["https://lb.example.test/session"]
|
| 535 |
+
assert requested_session_headers == [{"Authorization": "Bearer hf-secret"}]
|
| 536 |
+
assert captured_client_kwargs["api_key"] == "hf-secret"
|
| 537 |
+
assert captured_client_kwargs["base_url"] == "https://hf.example.test/v1"
|
| 538 |
+
assert captured_client_kwargs["websocket_base_url"] == "wss://hf.example.test/v1"
|
| 539 |
+
assert handler._realtime_connect_query == {"session_token": "allocated"}
|
| 540 |
+
|
| 541 |
+
|
| 542 |
+
@pytest.mark.asyncio
|
| 543 |
+
async def test_build_realtime_client_does_not_send_openai_key_to_hf_allocator(monkeypatch: Any) -> None:
|
| 544 |
+
"""Hugging Face allocator auth should use HF_TOKEN only."""
|
| 545 |
+
captured_client_kwargs: dict[str, Any] = {}
|
| 546 |
+
requested_session_headers: list[dict[str, str] | None] = []
|
| 547 |
+
|
| 548 |
+
class FakeClient:
|
| 549 |
+
def __init__(self, **kwargs: Any) -> None:
|
| 550 |
+
captured_client_kwargs.update(kwargs)
|
| 551 |
+
|
| 552 |
+
class FakeResponse:
|
| 553 |
+
def raise_for_status(self) -> None:
|
| 554 |
+
pass
|
| 555 |
+
|
| 556 |
+
def json(self) -> dict[str, str]:
|
| 557 |
+
return {
|
| 558 |
+
"session_id": "session-123",
|
| 559 |
+
"connect_url": "wss://hf.example.test/v1/realtime?session_token=allocated",
|
| 560 |
+
}
|
| 561 |
+
|
| 562 |
+
class FakeAsyncClient:
|
| 563 |
+
def __init__(self, **_kwargs: Any) -> None:
|
| 564 |
+
pass
|
| 565 |
+
|
| 566 |
+
async def __aenter__(self) -> "FakeAsyncClient":
|
| 567 |
+
return self
|
| 568 |
+
|
| 569 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 570 |
+
return False
|
| 571 |
+
|
| 572 |
+
async def post(self, _url: str, headers: dict[str, str] | None = None) -> FakeResponse:
|
| 573 |
+
requested_session_headers.append(headers)
|
| 574 |
+
return FakeResponse()
|
| 575 |
+
|
| 576 |
+
monkeypatch.setattr(hf_mod, "AsyncOpenAI", FakeClient)
|
| 577 |
+
monkeypatch.setattr(hf_mod.httpx, "AsyncClient", FakeAsyncClient)
|
| 578 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 579 |
+
monkeypatch.setattr(config, "HF_REALTIME_CONNECTION_MODE", "deployed")
|
| 580 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", "https://lb.example.test/session")
|
| 581 |
+
monkeypatch.setattr(config, "HF_REALTIME_WS_URL", None)
|
| 582 |
+
monkeypatch.setattr(config, "OPENAI_API_KEY", "sk-openai-secret")
|
| 583 |
+
monkeypatch.setattr(config, "HF_TOKEN", None)
|
| 584 |
+
|
| 585 |
+
handler = HuggingFaceRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 586 |
+
|
| 587 |
+
client = await handler._build_realtime_client()
|
| 588 |
+
|
| 589 |
+
assert client is not None
|
| 590 |
+
assert requested_session_headers == [None]
|
| 591 |
+
assert captured_client_kwargs["api_key"] == "DUMMY"
|
| 592 |
+
|
| 593 |
+
|
| 594 |
+
@pytest.mark.asyncio
|
| 595 |
+
async def test_apply_personality_uses_selected_voice_for_lb_allocated_sessions(monkeypatch: Any) -> None:
|
| 596 |
+
"""Live personality updates should honor the selected Qwen CustomVoice speaker."""
|
| 597 |
+
monkeypatch.setattr(hf_mod, "get_session_instructions", lambda: "new instructions")
|
| 598 |
+
monkeypatch.setattr(hf_mod, "get_session_voice", lambda default=HF_DEFAULT_VOICE: "Serena")
|
| 599 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "huggingface")
|
| 600 |
+
monkeypatch.setattr(config, "HF_REALTIME_SESSION_URL", "https://lb.example.test/session")
|
| 601 |
+
|
| 602 |
+
captured_update: dict[str, Any] = {}
|
| 603 |
+
|
| 604 |
+
class FakeSession:
|
| 605 |
+
async def update(self, **kwargs: Any) -> None:
|
| 606 |
+
captured_update.update(kwargs)
|
| 607 |
+
|
| 608 |
+
class FakeConnection:
|
| 609 |
+
session = FakeSession()
|
| 610 |
+
|
| 611 |
+
handler = HuggingFaceRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 612 |
+
handler.connection = FakeConnection()
|
| 613 |
+
monkeypatch.setattr(handler, "_restart_session", AsyncMock(return_value=None))
|
| 614 |
+
|
| 615 |
+
result = await handler.apply_personality("example")
|
| 616 |
+
|
| 617 |
+
assert "restarted realtime session" in result.lower()
|
| 618 |
+
session = captured_update["session"]
|
| 619 |
+
assert session["instructions"] == "new instructions"
|
| 620 |
+
assert session["audio"]["output"]["voice"] == "Serena"
|
| 621 |
+
|
| 622 |
+
|
| 623 |
+
def test_huggingface_response_cost_defaults_to_zero() -> None:
|
| 624 |
+
"""Hugging Face should not inherit OpenAI pricing from the shared base handler."""
|
| 625 |
+
usage = _make_usage(audio_in=1000, text_in=2000, image_in=500, audio_out=800, text_out=300)
|
| 626 |
+
handler = HuggingFaceRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 627 |
+
|
| 628 |
+
assert handler._compute_response_cost(usage) == 0.0
|
tests/test_openai_realtime.py
CHANGED
|
@@ -1,4 +1,3 @@
|
|
| 1 |
-
import base64
|
| 2 |
import random
|
| 3 |
import asyncio
|
| 4 |
import logging
|
|
@@ -8,26 +7,114 @@ from datetime import datetime, timezone
|
|
| 8 |
from unittest.mock import AsyncMock, MagicMock
|
| 9 |
|
| 10 |
import pytest
|
|
|
|
| 11 |
|
|
|
|
| 12 |
import reachy_mini_conversation_app.openai_realtime as rt_mod
|
|
|
|
| 13 |
import reachy_mini_conversation_app.tools.background_tool_manager as btm_mod
|
| 14 |
-
from reachy_mini_conversation_app.
|
|
|
|
| 15 |
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 16 |
from reachy_mini_conversation_app.tools.background_tool_manager import ToolCallRoutine
|
| 17 |
|
| 18 |
|
|
|
|
|
|
|
|
|
|
| 19 |
def _build_handler(loop: asyncio.AbstractEventLoop) -> OpenaiRealtimeHandler:
|
| 20 |
asyncio.set_event_loop(loop)
|
| 21 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 22 |
return OpenaiRealtimeHandler(deps)
|
| 23 |
|
| 24 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 25 |
@pytest.mark.asyncio
|
| 26 |
async def test_tool_completion_does_not_reset_head_wobbler(monkeypatch: Any) -> None:
|
| 27 |
"""Tool completion should not interrupt ongoing speech wobble."""
|
| 28 |
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 29 |
-
monkeypatch.setattr(rt_mod, "get_session_voice", lambda: "alloy")
|
| 30 |
-
monkeypatch.setattr(rt_mod, "
|
| 31 |
|
| 32 |
async def _fake_dispatch(tool_name: str, args_json: str, deps: Any, **_kw: Any) -> dict[str, Any]:
|
| 33 |
return {"image_description": "A person in front of a door.", "tool": tool_name}
|
|
@@ -104,7 +191,7 @@ async def test_tool_completion_does_not_reset_head_wobbler(monkeypatch: Any) ->
|
|
| 104 |
head_wobbler=head_wobbler,
|
| 105 |
)
|
| 106 |
handler = OpenaiRealtimeHandler(deps)
|
| 107 |
-
fake_client
|
| 108 |
handler.client = fake_client
|
| 109 |
|
| 110 |
session_task = asyncio.create_task(handler._run_realtime_session())
|
|
@@ -132,8 +219,8 @@ async def test_tool_completion_does_not_reset_head_wobbler(monkeypatch: Any) ->
|
|
| 132 |
async def test_non_idle_tool_call_does_not_queue_progress_response(monkeypatch: Any) -> None:
|
| 133 |
"""Tool-call startup should not enqueue a second speech response."""
|
| 134 |
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 135 |
-
monkeypatch.setattr(rt_mod, "get_session_voice", lambda: "alloy")
|
| 136 |
-
monkeypatch.setattr(rt_mod, "
|
| 137 |
|
| 138 |
class FakeEvent:
|
| 139 |
def __init__(self, etype: str, **kwargs: Any) -> None:
|
|
@@ -209,16 +296,16 @@ async def test_non_idle_tool_call_does_not_queue_progress_response(monkeypatch:
|
|
| 209 |
|
| 210 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 211 |
handler = OpenaiRealtimeHandler(deps)
|
| 212 |
-
fake_client
|
| 213 |
handler.client = fake_client
|
| 214 |
safe_response_create = AsyncMock()
|
| 215 |
-
|
| 216 |
start_up = MagicMock()
|
| 217 |
shutdown = AsyncMock()
|
| 218 |
start_tool = AsyncMock(return_value=MagicMock(tool_id="camera-call_camera_1-0"))
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
|
| 223 |
await handler._run_realtime_session()
|
| 224 |
|
|
@@ -227,11 +314,11 @@ async def test_non_idle_tool_call_does_not_queue_progress_response(monkeypatch:
|
|
| 227 |
|
| 228 |
|
| 229 |
@pytest.mark.asyncio
|
| 230 |
-
async def
|
| 231 |
-
"""
|
| 232 |
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 233 |
-
monkeypatch.setattr(rt_mod, "get_session_voice", lambda: "alloy")
|
| 234 |
-
monkeypatch.setattr(rt_mod, "
|
| 235 |
|
| 236 |
class FakeEvent:
|
| 237 |
def __init__(self, etype: str, **kwargs: Any) -> None:
|
|
@@ -270,12 +357,8 @@ async def test_output_audio_done_schedules_head_wobbler_reset(monkeypatch: Any)
|
|
| 270 |
def __init__(self) -> None:
|
| 271 |
self._events = iter(
|
| 272 |
[
|
| 273 |
-
FakeEvent("
|
| 274 |
-
FakeEvent(
|
| 275 |
-
"response.output_audio.delta",
|
| 276 |
-
delta=base64.b64encode(b"\x00\x00\x10\x00").decode("ascii"),
|
| 277 |
-
),
|
| 278 |
-
FakeEvent("response.output_audio.done"),
|
| 279 |
]
|
| 280 |
)
|
| 281 |
|
|
@@ -305,28 +388,121 @@ async def test_output_audio_done_schedules_head_wobbler_reset(monkeypatch: Any)
|
|
| 305 |
def __init__(self) -> None:
|
| 306 |
self.realtime = FakeRealtime()
|
| 307 |
|
| 308 |
-
|
| 309 |
-
|
| 310 |
-
|
| 311 |
-
|
| 312 |
-
|
| 313 |
-
|
| 314 |
-
head_wobbler=head_wobbler,
|
| 315 |
-
)
|
| 316 |
-
handler = OpenaiRealtimeHandler(deps, gradio_mode=True)
|
| 317 |
-
handler.client = FakeClient()
|
| 318 |
-
object.__setattr__(handler.tool_manager, "start_up", MagicMock())
|
| 319 |
-
object.__setattr__(handler.tool_manager, "shutdown", AsyncMock())
|
| 320 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 321 |
await handler._run_realtime_session()
|
| 322 |
|
| 323 |
-
|
| 324 |
-
|
| 325 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 326 |
|
| 327 |
|
| 328 |
def test_format_timestamp_uses_wall_clock() -> None:
|
| 329 |
"""Test that format_timestamp uses wall clock time."""
|
|
|
|
|
|
|
|
|
|
|
|
|
| 330 |
loop = asyncio.new_event_loop()
|
| 331 |
try:
|
| 332 |
print("Testing format_timestamp...")
|
|
@@ -334,8 +510,8 @@ def test_format_timestamp_uses_wall_clock() -> None:
|
|
| 334 |
formatted = handler.format_timestamp()
|
| 335 |
print(f"Formatted timestamp: {formatted}")
|
| 336 |
finally:
|
| 337 |
-
asyncio.set_event_loop(None)
|
| 338 |
loop.close()
|
|
|
|
| 339 |
|
| 340 |
# Extract year from "[YYYY-MM-DD ...]"
|
| 341 |
year = int(formatted[1:5])
|
|
@@ -351,9 +527,9 @@ async def test_start_up_retries_on_abrupt_close(monkeypatch: Any, caplog: Any) -
|
|
| 351 |
"""
|
| 352 |
caplog.set_level(logging.WARNING)
|
| 353 |
|
| 354 |
-
# Use a local Exception as the module's ConnectionClosedError to avoid ws dependency
|
| 355 |
FakeCCE = type("FakeCCE", (Exception,), {})
|
| 356 |
-
monkeypatch.setattr(
|
| 357 |
|
| 358 |
# Make asyncio.sleep return immediately (for backoff)
|
| 359 |
_real_sleep = asyncio.sleep
|
|
@@ -431,6 +607,7 @@ async def test_start_up_retries_on_abrupt_close(monkeypatch: Any, caplog: Any) -
|
|
| 431 |
|
| 432 |
# Patch the OpenAI client used by the handler
|
| 433 |
monkeypatch.setattr(rt_mod, "AsyncOpenAI", FakeClient)
|
|
|
|
| 434 |
|
| 435 |
# Build handler with minimal deps
|
| 436 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
|
@@ -448,6 +625,78 @@ async def test_start_up_retries_on_abrupt_close(monkeypatch: Any, caplog: Any) -
|
|
| 448 |
assert len(warnings) == 1
|
| 449 |
|
| 450 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 451 |
# ---- Cost calculation tests ----
|
| 452 |
|
| 453 |
|
|
@@ -495,9 +744,10 @@ def _make_usage(
|
|
| 495 |
ids=["normal", "all_none", "mixed", "missing_details"],
|
| 496 |
)
|
| 497 |
def test_compute_response_cost(usage_kwargs: dict[str, Any], expect_positive: bool) -> None:
|
| 498 |
-
"""Verify
|
| 499 |
usage = _make_usage(**usage_kwargs)
|
| 500 |
-
|
|
|
|
| 501 |
if expect_positive:
|
| 502 |
assert cost > 0
|
| 503 |
else:
|
|
@@ -507,6 +757,138 @@ def test_compute_response_cost(usage_kwargs: dict[str, Any], expect_positive: bo
|
|
| 507 |
# ---- Stress test: response.create rejection + retry ----
|
| 508 |
|
| 509 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 510 |
@pytest.mark.asyncio
|
| 511 |
async def test_response_sender_retries_on_active_response_rejection(monkeypatch: Any, caplog: Any) -> None:
|
| 512 |
"""Stress test: response.create rejection + retry via real event processing.
|
|
@@ -521,20 +903,20 @@ async def test_response_sender_retries_on_active_response_rejection(monkeypatch:
|
|
| 521 |
processing, not mocked out.
|
| 522 |
"""
|
| 523 |
caplog.set_level(logging.DEBUG)
|
|
|
|
| 524 |
|
| 525 |
FakeCCE = type("FakeCCE", (Exception,), {})
|
| 526 |
-
monkeypatch.setattr(
|
| 527 |
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 528 |
-
monkeypatch.setattr(rt_mod, "get_session_voice", lambda: "alloy")
|
| 529 |
-
monkeypatch.setattr(rt_mod, "
|
| 530 |
|
| 531 |
N_TOOL_RESULTS = 400
|
| 532 |
REJECT_CALL_NUMBERS = {1, 3, 5, 10, 25, 50, 75, 100, 150, 200, 300, 399}
|
| 533 |
EXPECTED_TOTAL_CALLS = N_TOOL_RESULTS + len(REJECT_CALL_NUMBERS)
|
| 534 |
|
| 535 |
-
event_queue: asyncio.Queue[Any] = asyncio.Queue()
|
| 536 |
response_create_log: list[tuple[int, dict[str, Any]]] = []
|
| 537 |
-
handler_ref: list[
|
| 538 |
|
| 539 |
# ---- Fake event / error objects mirroring the OpenAI SDK shapes ----
|
| 540 |
|
|
@@ -558,6 +940,8 @@ async def test_response_sender_retries_on_active_response_rejection(monkeypatch:
|
|
| 558 |
for k, v in kwargs.items():
|
| 559 |
setattr(self, k, v)
|
| 560 |
|
|
|
|
|
|
|
| 561 |
# ---- Fake connection components ----
|
| 562 |
|
| 563 |
class FakeResponseAPI:
|
|
@@ -660,7 +1044,7 @@ async def test_response_sender_retries_on_active_response_rejection(monkeypatch:
|
|
| 660 |
return self
|
| 661 |
|
| 662 |
async def __anext__(self) -> FakeEvent:
|
| 663 |
-
event
|
| 664 |
if event is None: # sentinel → end iteration
|
| 665 |
raise StopAsyncIteration
|
| 666 |
return event
|
|
@@ -674,6 +1058,7 @@ async def test_response_sender_retries_on_active_response_rejection(monkeypatch:
|
|
| 674 |
self.realtime = FakeRealtime()
|
| 675 |
|
| 676 |
monkeypatch.setattr(rt_mod, "AsyncOpenAI", FakeClient)
|
|
|
|
| 677 |
|
| 678 |
# Patch dispatch_tool_call so tools complete with a result.
|
| 679 |
async def _fake_dispatch(tool_name: str, args_json: str, deps: Any, **_kw: Any) -> dict[str, Any]:
|
|
@@ -704,10 +1089,13 @@ async def test_response_sender_retries_on_active_response_rejection(monkeypatch:
|
|
| 704 |
is_idle_tool_call=False,
|
| 705 |
)
|
| 706 |
|
| 707 |
-
#
|
| 708 |
-
# This stress test queues hundreds of serialized response.create calls
|
| 709 |
-
#
|
| 710 |
-
|
|
|
|
|
|
|
|
|
|
| 711 |
|
| 712 |
# ---- Tear down ----
|
| 713 |
|
|
@@ -760,7 +1148,7 @@ async def test_response_sender_loop_times_out_waiting_for_response_done(
|
|
| 760 |
"""
|
| 761 |
caplog.set_level(logging.DEBUG)
|
| 762 |
|
| 763 |
-
monkeypatch.setattr(
|
| 764 |
|
| 765 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 766 |
handler = rt_mod.OpenaiRealtimeHandler(deps)
|
|
@@ -812,7 +1200,7 @@ async def test_response_sender_loop_times_out_waiting_for_previous_response(
|
|
| 812 |
"""
|
| 813 |
caplog.set_level(logging.DEBUG)
|
| 814 |
|
| 815 |
-
monkeypatch.setattr(
|
| 816 |
|
| 817 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 818 |
handler = rt_mod.OpenaiRealtimeHandler(deps)
|
|
@@ -848,3 +1236,104 @@ async def test_response_sender_loop_times_out_waiting_for_previous_response(
|
|
| 848 |
|
| 849 |
timeout_logs = [r for r in caplog.records if "Timed out waiting for previous response" in r.getMessage()]
|
| 850 |
assert len(timeout_logs) == 1, f"Expected 1 pre-condition timeout warning, got {len(timeout_logs)}"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
import random
|
| 2 |
import asyncio
|
| 3 |
import logging
|
|
|
|
| 7 |
from unittest.mock import AsyncMock, MagicMock
|
| 8 |
|
| 9 |
import pytest
|
| 10 |
+
from fastrtc import AdditionalOutputs
|
| 11 |
|
| 12 |
+
import reachy_mini_conversation_app.base_realtime as base_rt_mod
|
| 13 |
import reachy_mini_conversation_app.openai_realtime as rt_mod
|
| 14 |
+
import reachy_mini_conversation_app.tools.core_tools as ct_mod
|
| 15 |
import reachy_mini_conversation_app.tools.background_tool_manager as btm_mod
|
| 16 |
+
from reachy_mini_conversation_app.config import OPENAI_BACKEND, config, get_default_voice_for_backend
|
| 17 |
+
from reachy_mini_conversation_app.openai_realtime import OpenaiRealtimeHandler
|
| 18 |
from reachy_mini_conversation_app.tools.core_tools import ToolDependencies
|
| 19 |
from reachy_mini_conversation_app.tools.background_tool_manager import ToolCallRoutine
|
| 20 |
|
| 21 |
|
| 22 |
+
OPENAI_DEFAULT_VOICE = get_default_voice_for_backend(OPENAI_BACKEND)
|
| 23 |
+
|
| 24 |
+
|
| 25 |
def _build_handler(loop: asyncio.AbstractEventLoop) -> OpenaiRealtimeHandler:
|
| 26 |
asyncio.set_event_loop(loop)
|
| 27 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 28 |
return OpenaiRealtimeHandler(deps)
|
| 29 |
|
| 30 |
|
| 31 |
+
async def _run_openai_handler_with_events(
|
| 32 |
+
monkeypatch: Any,
|
| 33 |
+
events: list[Any],
|
| 34 |
+
*,
|
| 35 |
+
movement_manager: MagicMock | None = None,
|
| 36 |
+
) -> OpenaiRealtimeHandler:
|
| 37 |
+
"""Run an OpenAI realtime handler against a fixed event sequence."""
|
| 38 |
+
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 39 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda default=OPENAI_DEFAULT_VOICE: "alloy")
|
| 40 |
+
monkeypatch.setattr(rt_mod, "get_active_tool_specs", lambda _: [])
|
| 41 |
+
|
| 42 |
+
class FakeSession:
|
| 43 |
+
async def update(self, **_kw: Any) -> None:
|
| 44 |
+
pass
|
| 45 |
+
|
| 46 |
+
class FakeInputAudioBuffer:
|
| 47 |
+
async def append(self, **_kw: Any) -> None:
|
| 48 |
+
pass
|
| 49 |
+
|
| 50 |
+
class FakeItem:
|
| 51 |
+
async def create(self, **_kw: Any) -> None:
|
| 52 |
+
pass
|
| 53 |
+
|
| 54 |
+
class FakeConversation:
|
| 55 |
+
item = FakeItem()
|
| 56 |
+
|
| 57 |
+
class FakeResponse:
|
| 58 |
+
async def create(self, **_kw: Any) -> None:
|
| 59 |
+
pass
|
| 60 |
+
|
| 61 |
+
async def cancel(self, **_kw: Any) -> None:
|
| 62 |
+
pass
|
| 63 |
+
|
| 64 |
+
class FakeConn:
|
| 65 |
+
session = FakeSession()
|
| 66 |
+
input_audio_buffer = FakeInputAudioBuffer()
|
| 67 |
+
conversation = FakeConversation()
|
| 68 |
+
response = FakeResponse()
|
| 69 |
+
|
| 70 |
+
def __init__(self) -> None:
|
| 71 |
+
self._events = iter(events)
|
| 72 |
+
|
| 73 |
+
async def __aenter__(self) -> "FakeConn":
|
| 74 |
+
return self
|
| 75 |
+
|
| 76 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 77 |
+
return False
|
| 78 |
+
|
| 79 |
+
async def close(self) -> None:
|
| 80 |
+
pass
|
| 81 |
+
|
| 82 |
+
def __aiter__(self) -> "FakeConn":
|
| 83 |
+
return self
|
| 84 |
+
|
| 85 |
+
async def __anext__(self) -> Any:
|
| 86 |
+
try:
|
| 87 |
+
return next(self._events)
|
| 88 |
+
except StopIteration:
|
| 89 |
+
raise StopAsyncIteration
|
| 90 |
+
|
| 91 |
+
class FakeRealtime:
|
| 92 |
+
def connect(self, **_kw: Any) -> FakeConn:
|
| 93 |
+
return FakeConn()
|
| 94 |
+
|
| 95 |
+
class FakeClient:
|
| 96 |
+
def __init__(self) -> None:
|
| 97 |
+
self.realtime = FakeRealtime()
|
| 98 |
+
|
| 99 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=movement_manager or MagicMock())
|
| 100 |
+
handler = OpenaiRealtimeHandler(deps)
|
| 101 |
+
handler.client = FakeClient()
|
| 102 |
+
|
| 103 |
+
start_up = MagicMock()
|
| 104 |
+
shutdown = AsyncMock()
|
| 105 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_up", start_up)
|
| 106 |
+
monkeypatch.setattr(type(handler.tool_manager), "shutdown", shutdown)
|
| 107 |
+
|
| 108 |
+
await handler._run_realtime_session()
|
| 109 |
+
return handler
|
| 110 |
+
|
| 111 |
+
|
| 112 |
@pytest.mark.asyncio
|
| 113 |
async def test_tool_completion_does_not_reset_head_wobbler(monkeypatch: Any) -> None:
|
| 114 |
"""Tool completion should not interrupt ongoing speech wobble."""
|
| 115 |
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 116 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda default=OPENAI_DEFAULT_VOICE: "alloy")
|
| 117 |
+
monkeypatch.setattr(rt_mod, "get_active_tool_specs", lambda _: [])
|
| 118 |
|
| 119 |
async def _fake_dispatch(tool_name: str, args_json: str, deps: Any, **_kw: Any) -> dict[str, Any]:
|
| 120 |
return {"image_description": "A person in front of a door.", "tool": tool_name}
|
|
|
|
| 191 |
head_wobbler=head_wobbler,
|
| 192 |
)
|
| 193 |
handler = OpenaiRealtimeHandler(deps)
|
| 194 |
+
fake_client = FakeClient()
|
| 195 |
handler.client = fake_client
|
| 196 |
|
| 197 |
session_task = asyncio.create_task(handler._run_realtime_session())
|
|
|
|
| 219 |
async def test_non_idle_tool_call_does_not_queue_progress_response(monkeypatch: Any) -> None:
|
| 220 |
"""Tool-call startup should not enqueue a second speech response."""
|
| 221 |
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 222 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda default=OPENAI_DEFAULT_VOICE: "alloy")
|
| 223 |
+
monkeypatch.setattr(rt_mod, "get_active_tool_specs", lambda _: [])
|
| 224 |
|
| 225 |
class FakeEvent:
|
| 226 |
def __init__(self, etype: str, **kwargs: Any) -> None:
|
|
|
|
| 296 |
|
| 297 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 298 |
handler = OpenaiRealtimeHandler(deps)
|
| 299 |
+
fake_client = FakeClient()
|
| 300 |
handler.client = fake_client
|
| 301 |
safe_response_create = AsyncMock()
|
| 302 |
+
monkeypatch.setattr(handler, "_safe_response_create", safe_response_create)
|
| 303 |
start_up = MagicMock()
|
| 304 |
shutdown = AsyncMock()
|
| 305 |
start_tool = AsyncMock(return_value=MagicMock(tool_id="camera-call_camera_1-0"))
|
| 306 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_up", start_up)
|
| 307 |
+
monkeypatch.setattr(type(handler.tool_manager), "shutdown", shutdown)
|
| 308 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_tool", start_tool)
|
| 309 |
|
| 310 |
await handler._run_realtime_session()
|
| 311 |
|
|
|
|
| 314 |
|
| 315 |
|
| 316 |
@pytest.mark.asyncio
|
| 317 |
+
async def test_user_speech_events_reset_idle_timer(monkeypatch: Any) -> None:
|
| 318 |
+
"""User speech/transcription events should postpone idle behavior."""
|
| 319 |
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 320 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda default=OPENAI_DEFAULT_VOICE: "alloy")
|
| 321 |
+
monkeypatch.setattr(rt_mod, "get_active_tool_specs", lambda _: [])
|
| 322 |
|
| 323 |
class FakeEvent:
|
| 324 |
def __init__(self, etype: str, **kwargs: Any) -> None:
|
|
|
|
| 357 |
def __init__(self) -> None:
|
| 358 |
self._events = iter(
|
| 359 |
[
|
| 360 |
+
FakeEvent("input_audio_buffer.speech_started"),
|
| 361 |
+
FakeEvent("conversation.item.input_audio_transcription.completed", transcript="hello there"),
|
|
|
|
|
|
|
|
|
|
|
|
|
| 362 |
]
|
| 363 |
)
|
| 364 |
|
|
|
|
| 388 |
def __init__(self) -> None:
|
| 389 |
self.realtime = FakeRealtime()
|
| 390 |
|
| 391 |
+
movement_manager = MagicMock()
|
| 392 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=movement_manager)
|
| 393 |
+
handler = OpenaiRealtimeHandler(deps)
|
| 394 |
+
fake_client = FakeClient()
|
| 395 |
+
handler.client = fake_client
|
| 396 |
+
handler.last_activity_time = asyncio.get_running_loop().time() - 60.0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 397 |
|
| 398 |
+
start_up = MagicMock()
|
| 399 |
+
shutdown = AsyncMock()
|
| 400 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_up", start_up)
|
| 401 |
+
monkeypatch.setattr(type(handler.tool_manager), "shutdown", shutdown)
|
| 402 |
+
|
| 403 |
+
previous_activity_time = handler.last_activity_time
|
| 404 |
await handler._run_realtime_session()
|
| 405 |
|
| 406 |
+
assert handler.last_activity_time > previous_activity_time
|
| 407 |
+
movement_manager.set_listening.assert_any_call(True)
|
| 408 |
+
|
| 409 |
+
|
| 410 |
+
@pytest.mark.asyncio
|
| 411 |
+
async def test_empty_user_transcript_exits_listening_without_chat_message(monkeypatch: Any) -> None:
|
| 412 |
+
"""Blank VAD commits should not leave listening motion frozen."""
|
| 413 |
+
movement_manager = MagicMock()
|
| 414 |
+
|
| 415 |
+
handler = await _run_openai_handler_with_events(
|
| 416 |
+
monkeypatch,
|
| 417 |
+
[
|
| 418 |
+
SimpleNamespace(type="input_audio_buffer.speech_started"),
|
| 419 |
+
SimpleNamespace(type="conversation.item.input_audio_transcription.completed", transcript=" "),
|
| 420 |
+
],
|
| 421 |
+
movement_manager=movement_manager,
|
| 422 |
+
)
|
| 423 |
+
|
| 424 |
+
assert [call.args[0] for call in movement_manager.set_listening.call_args_list] == [True, False]
|
| 425 |
+
assert handler.output_queue.empty()
|
| 426 |
+
assert handler._turn_user_done_at is None
|
| 427 |
+
|
| 428 |
+
|
| 429 |
+
@pytest.mark.asyncio
|
| 430 |
+
async def test_empty_audio_buffer_error_exits_listening_without_chat_error(monkeypatch: Any) -> None:
|
| 431 |
+
"""Empty audio-buffer commits are internal and should restore listening state."""
|
| 432 |
+
movement_manager = MagicMock()
|
| 433 |
+
|
| 434 |
+
handler = await _run_openai_handler_with_events(
|
| 435 |
+
monkeypatch,
|
| 436 |
+
[
|
| 437 |
+
SimpleNamespace(type="input_audio_buffer.speech_started"),
|
| 438 |
+
SimpleNamespace(
|
| 439 |
+
type="error",
|
| 440 |
+
error=SimpleNamespace(code="input_audio_buffer_commit_empty", message="empty audio buffer"),
|
| 441 |
+
),
|
| 442 |
+
],
|
| 443 |
+
movement_manager=movement_manager,
|
| 444 |
+
)
|
| 445 |
+
|
| 446 |
+
assert [call.args[0] for call in movement_manager.set_listening.call_args_list] == [True, False]
|
| 447 |
+
assert handler.output_queue.empty()
|
| 448 |
+
|
| 449 |
+
|
| 450 |
+
@pytest.mark.asyncio
|
| 451 |
+
async def test_apply_personality_preserves_manual_voice_override(monkeypatch: Any) -> None:
|
| 452 |
+
"""Applying a profile should not discard a voice manually selected in the current session."""
|
| 453 |
+
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 454 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda: "cedar")
|
| 455 |
+
monkeypatch.setattr("reachy_mini_conversation_app.config.set_custom_profile", lambda _profile: None)
|
| 456 |
+
|
| 457 |
+
handler = OpenaiRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 458 |
+
update = AsyncMock()
|
| 459 |
+
handler.connection = SimpleNamespace(session=SimpleNamespace(update=update))
|
| 460 |
+
handler._voice_override = "marin"
|
| 461 |
+
restart = AsyncMock()
|
| 462 |
+
monkeypatch.setattr(handler, "_restart_session", restart)
|
| 463 |
+
|
| 464 |
+
status = await handler.apply_personality("example")
|
| 465 |
+
|
| 466 |
+
assert status == "Applied personality and restarted realtime session."
|
| 467 |
+
assert handler.get_current_voice() == "marin"
|
| 468 |
+
restart.assert_awaited_once()
|
| 469 |
+
session = update.await_args.kwargs["session"]
|
| 470 |
+
assert session["audio"]["output"]["voice"] == "marin"
|
| 471 |
+
|
| 472 |
+
|
| 473 |
+
def test_handler_uses_startup_voice_at_startup(monkeypatch: Any) -> None:
|
| 474 |
+
"""OpenAI handler startup should restore a persisted startup voice."""
|
| 475 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 476 |
+
|
| 477 |
+
handler = OpenaiRealtimeHandler(
|
| 478 |
+
ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()),
|
| 479 |
+
startup_voice="shimmer",
|
| 480 |
+
)
|
| 481 |
+
|
| 482 |
+
assert handler.get_current_voice() == "shimmer"
|
| 483 |
+
|
| 484 |
+
|
| 485 |
+
def test_copy_preserves_current_voice_override(monkeypatch: Any) -> None:
|
| 486 |
+
"""Copied OpenAI handlers should keep the current voice override."""
|
| 487 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 488 |
+
|
| 489 |
+
handler = OpenaiRealtimeHandler(
|
| 490 |
+
ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()),
|
| 491 |
+
startup_voice="shimmer",
|
| 492 |
+
)
|
| 493 |
+
handler._voice_override = "marin"
|
| 494 |
+
|
| 495 |
+
copied_handler = handler.copy()
|
| 496 |
+
|
| 497 |
+
assert copied_handler.get_current_voice() == "marin"
|
| 498 |
|
| 499 |
|
| 500 |
def test_format_timestamp_uses_wall_clock() -> None:
|
| 501 |
"""Test that format_timestamp uses wall clock time."""
|
| 502 |
+
try:
|
| 503 |
+
previous_loop = asyncio.get_event_loop()
|
| 504 |
+
except RuntimeError:
|
| 505 |
+
previous_loop = asyncio.new_event_loop()
|
| 506 |
loop = asyncio.new_event_loop()
|
| 507 |
try:
|
| 508 |
print("Testing format_timestamp...")
|
|
|
|
| 510 |
formatted = handler.format_timestamp()
|
| 511 |
print(f"Formatted timestamp: {formatted}")
|
| 512 |
finally:
|
|
|
|
| 513 |
loop.close()
|
| 514 |
+
asyncio.set_event_loop(previous_loop)
|
| 515 |
|
| 516 |
# Extract year from "[YYYY-MM-DD ...]"
|
| 517 |
year = int(formatted[1:5])
|
|
|
|
| 527 |
"""
|
| 528 |
caplog.set_level(logging.WARNING)
|
| 529 |
|
| 530 |
+
# Use a local Exception as the base module's ConnectionClosedError to avoid ws dependency.
|
| 531 |
FakeCCE = type("FakeCCE", (Exception,), {})
|
| 532 |
+
monkeypatch.setattr(base_rt_mod, "ConnectionClosedError", FakeCCE)
|
| 533 |
|
| 534 |
# Make asyncio.sleep return immediately (for backoff)
|
| 535 |
_real_sleep = asyncio.sleep
|
|
|
|
| 607 |
|
| 608 |
# Patch the OpenAI client used by the handler
|
| 609 |
monkeypatch.setattr(rt_mod, "AsyncOpenAI", FakeClient)
|
| 610 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 611 |
|
| 612 |
# Build handler with minimal deps
|
| 613 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
|
|
|
| 625 |
assert len(warnings) == 1
|
| 626 |
|
| 627 |
|
| 628 |
+
@pytest.mark.asyncio
|
| 629 |
+
async def test_start_up_openai_gradio_collects_textbox_api_key(monkeypatch: Any) -> None:
|
| 630 |
+
"""OpenAI should own Gradio textbox credential collection."""
|
| 631 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 632 |
+
monkeypatch.setattr(config, "OPENAI_API_KEY", None)
|
| 633 |
+
|
| 634 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 635 |
+
handler = rt_mod.OpenaiRealtimeHandler(deps, gradio_mode=True)
|
| 636 |
+
handler.latest_args = ["profile", "voice", "unused", "sk-textbox-secret"]
|
| 637 |
+
|
| 638 |
+
build_client = AsyncMock(return_value=MagicMock())
|
| 639 |
+
run_realtime_session = AsyncMock(return_value=None)
|
| 640 |
+
wait_for_args = AsyncMock(return_value=None)
|
| 641 |
+
|
| 642 |
+
monkeypatch.setattr(handler, "_build_realtime_client", build_client)
|
| 643 |
+
monkeypatch.setattr(handler, "_run_realtime_session", run_realtime_session)
|
| 644 |
+
monkeypatch.setattr(handler, "wait_for_args", wait_for_args)
|
| 645 |
+
|
| 646 |
+
await handler.start_up()
|
| 647 |
+
|
| 648 |
+
wait_for_args.assert_awaited_once()
|
| 649 |
+
build_client.assert_awaited_once_with()
|
| 650 |
+
run_realtime_session.assert_awaited_once()
|
| 651 |
+
assert handler._provided_api_key == "sk-textbox-secret"
|
| 652 |
+
|
| 653 |
+
|
| 654 |
+
@pytest.mark.asyncio
|
| 655 |
+
async def test_run_realtime_session_propagates_session_update_failure(monkeypatch: Any) -> None:
|
| 656 |
+
"""A failed session.update must abort startup instead of looking like a clean session exit."""
|
| 657 |
+
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 658 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda default=OPENAI_DEFAULT_VOICE: "alloy")
|
| 659 |
+
monkeypatch.setattr(rt_mod, "get_active_tool_specs", lambda _: [])
|
| 660 |
+
|
| 661 |
+
class FakeSession:
|
| 662 |
+
async def update(self, **_kw: Any) -> None:
|
| 663 |
+
raise RuntimeError("invalid session config")
|
| 664 |
+
|
| 665 |
+
class FakeConn:
|
| 666 |
+
session = FakeSession()
|
| 667 |
+
|
| 668 |
+
async def __aenter__(self) -> "FakeConn":
|
| 669 |
+
return self
|
| 670 |
+
|
| 671 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 672 |
+
return False
|
| 673 |
+
|
| 674 |
+
class FakeRealtime:
|
| 675 |
+
def connect(self, **_kw: Any) -> FakeConn:
|
| 676 |
+
return FakeConn()
|
| 677 |
+
|
| 678 |
+
class FakeClient:
|
| 679 |
+
def __init__(self) -> None:
|
| 680 |
+
self.realtime = FakeRealtime()
|
| 681 |
+
|
| 682 |
+
handler = rt_mod.OpenaiRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 683 |
+
handler.client = FakeClient()
|
| 684 |
+
|
| 685 |
+
with pytest.raises(RuntimeError, match="invalid session config"):
|
| 686 |
+
await handler._run_realtime_session()
|
| 687 |
+
|
| 688 |
+
|
| 689 |
+
@pytest.mark.asyncio
|
| 690 |
+
async def test_handler_uses_openai_sample_rate_for_openai_backend(monkeypatch: Any) -> None:
|
| 691 |
+
"""OpenAI backend should keep the 24 kHz realtime audio configuration."""
|
| 692 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 693 |
+
|
| 694 |
+
handler = OpenaiRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 695 |
+
|
| 696 |
+
assert handler.input_sample_rate == 24000
|
| 697 |
+
assert handler.output_sample_rate == 24000
|
| 698 |
+
|
| 699 |
+
|
| 700 |
# ---- Cost calculation tests ----
|
| 701 |
|
| 702 |
|
|
|
|
| 744 |
ids=["normal", "all_none", "mixed", "missing_details"],
|
| 745 |
)
|
| 746 |
def test_compute_response_cost(usage_kwargs: dict[str, Any], expect_positive: bool) -> None:
|
| 747 |
+
"""Verify handler cost computation handles various token combinations without crashing."""
|
| 748 |
usage = _make_usage(**usage_kwargs)
|
| 749 |
+
handler = OpenaiRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 750 |
+
cost = handler._compute_response_cost(usage)
|
| 751 |
if expect_positive:
|
| 752 |
assert cost > 0
|
| 753 |
else:
|
|
|
|
| 757 |
# ---- Stress test: response.create rejection + retry ----
|
| 758 |
|
| 759 |
|
| 760 |
+
@pytest.mark.asyncio
|
| 761 |
+
async def test_response_sender_retries_when_active_response_error_uses_type_only(
|
| 762 |
+
monkeypatch: Any,
|
| 763 |
+
caplog: Any,
|
| 764 |
+
) -> None:
|
| 765 |
+
"""Retry active-response rejections even when the server omits ``error.code``.
|
| 766 |
+
|
| 767 |
+
Some backends only populate ``error.type=conversation_already_has_active_response``.
|
| 768 |
+
That should still take the retry path and must not be surfaced as a user-facing error.
|
| 769 |
+
"""
|
| 770 |
+
caplog.set_level(logging.DEBUG)
|
| 771 |
+
monkeypatch.setattr(base_rt_mod, "_RESPONSE_REJECTION_RETRY_DELAY", 0.01)
|
| 772 |
+
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 773 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda default=OPENAI_DEFAULT_VOICE: "alloy")
|
| 774 |
+
monkeypatch.setattr(rt_mod, "get_active_tool_specs", lambda _: [])
|
| 775 |
+
|
| 776 |
+
class FakeError:
|
| 777 |
+
def __init__(self, message: str) -> None:
|
| 778 |
+
self.message = message
|
| 779 |
+
self.code = None
|
| 780 |
+
self.type = "conversation_already_has_active_response"
|
| 781 |
+
self.event_id = None
|
| 782 |
+
self.param = None
|
| 783 |
+
|
| 784 |
+
def __repr__(self) -> str:
|
| 785 |
+
return f"RealtimeError(message='{self.message}', type='{self.type}', code=None, event_id=None, param=None)"
|
| 786 |
+
|
| 787 |
+
class FakeEvent:
|
| 788 |
+
def __init__(self, etype: str, **kwargs: Any) -> None:
|
| 789 |
+
self.type = etype
|
| 790 |
+
for key, value in kwargs.items():
|
| 791 |
+
setattr(self, key, value)
|
| 792 |
+
|
| 793 |
+
event_queue: asyncio.Queue[FakeEvent | None] = asyncio.Queue()
|
| 794 |
+
|
| 795 |
+
class FakeSession:
|
| 796 |
+
async def update(self, **_kw: Any) -> None:
|
| 797 |
+
pass
|
| 798 |
+
|
| 799 |
+
class FakeInputAudioBuffer:
|
| 800 |
+
async def append(self, **_kw: Any) -> None:
|
| 801 |
+
pass
|
| 802 |
+
|
| 803 |
+
class FakeItem:
|
| 804 |
+
async def create(self, **_kw: Any) -> None:
|
| 805 |
+
pass
|
| 806 |
+
|
| 807 |
+
class FakeConversation:
|
| 808 |
+
item = FakeItem()
|
| 809 |
+
|
| 810 |
+
class FakeResponseAPI:
|
| 811 |
+
def __init__(self) -> None:
|
| 812 |
+
self.call_count = 0
|
| 813 |
+
|
| 814 |
+
async def create(self, **_kw: Any) -> None:
|
| 815 |
+
self.call_count += 1
|
| 816 |
+
if self.call_count == 1:
|
| 817 |
+
event_queue.put_nowait(
|
| 818 |
+
FakeEvent(
|
| 819 |
+
"error",
|
| 820 |
+
error=FakeError("Cannot create response while another response is in progress."),
|
| 821 |
+
)
|
| 822 |
+
)
|
| 823 |
+
else:
|
| 824 |
+
event_queue.put_nowait(FakeEvent("response.created"))
|
| 825 |
+
event_queue.put_nowait(FakeEvent("response.done", response=MagicMock()))
|
| 826 |
+
|
| 827 |
+
async def cancel(self, **_kw: Any) -> None:
|
| 828 |
+
pass
|
| 829 |
+
|
| 830 |
+
fake_response_api = FakeResponseAPI()
|
| 831 |
+
|
| 832 |
+
class FakeConn:
|
| 833 |
+
session = FakeSession()
|
| 834 |
+
input_audio_buffer = FakeInputAudioBuffer()
|
| 835 |
+
conversation = FakeConversation()
|
| 836 |
+
response = fake_response_api
|
| 837 |
+
|
| 838 |
+
async def __aenter__(self) -> "FakeConn":
|
| 839 |
+
return self
|
| 840 |
+
|
| 841 |
+
async def __aexit__(self, *_args: Any) -> bool:
|
| 842 |
+
return False
|
| 843 |
+
|
| 844 |
+
async def close(self) -> None:
|
| 845 |
+
pass
|
| 846 |
+
|
| 847 |
+
def __aiter__(self) -> "FakeConn":
|
| 848 |
+
return self
|
| 849 |
+
|
| 850 |
+
async def __anext__(self) -> FakeEvent:
|
| 851 |
+
event = await event_queue.get()
|
| 852 |
+
if event is None:
|
| 853 |
+
raise StopAsyncIteration
|
| 854 |
+
return event
|
| 855 |
+
|
| 856 |
+
class FakeRealtime:
|
| 857 |
+
def connect(self, **_kw: Any) -> FakeConn:
|
| 858 |
+
return FakeConn()
|
| 859 |
+
|
| 860 |
+
class FakeClient:
|
| 861 |
+
def __init__(self) -> None:
|
| 862 |
+
self.realtime = FakeRealtime()
|
| 863 |
+
|
| 864 |
+
handler = rt_mod.OpenaiRealtimeHandler(ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock()))
|
| 865 |
+
handler.client = FakeClient()
|
| 866 |
+
|
| 867 |
+
session_task = asyncio.create_task(handler._run_realtime_session())
|
| 868 |
+
await asyncio.sleep(0)
|
| 869 |
+
await handler._safe_response_create(instructions="req")
|
| 870 |
+
await asyncio.sleep(0.1)
|
| 871 |
+
await event_queue.put(None)
|
| 872 |
+
await asyncio.wait_for(session_task, timeout=2.0)
|
| 873 |
+
|
| 874 |
+
assert fake_response_api.call_count == 2
|
| 875 |
+
assert not any(
|
| 876 |
+
record.levelname == "ERROR" and "Realtime error" in record.getMessage() for record in caplog.records
|
| 877 |
+
)
|
| 878 |
+
assert any("worker will retry after active response finishes" in record.getMessage() for record in caplog.records)
|
| 879 |
+
queued_outputs = []
|
| 880 |
+
while not handler.output_queue.empty():
|
| 881 |
+
queued_outputs.append(handler.output_queue.get_nowait())
|
| 882 |
+
queued_messages = [
|
| 883 |
+
message
|
| 884 |
+
for output in queued_outputs
|
| 885 |
+
if isinstance(output, AdditionalOutputs)
|
| 886 |
+
for message in output.args
|
| 887 |
+
if isinstance(message, dict)
|
| 888 |
+
]
|
| 889 |
+
assert not any(str(message.get("content", "")).startswith("[error]") for message in queued_messages)
|
| 890 |
+
|
| 891 |
+
|
| 892 |
@pytest.mark.asyncio
|
| 893 |
async def test_response_sender_retries_on_active_response_rejection(monkeypatch: Any, caplog: Any) -> None:
|
| 894 |
"""Stress test: response.create rejection + retry via real event processing.
|
|
|
|
| 903 |
processing, not mocked out.
|
| 904 |
"""
|
| 905 |
caplog.set_level(logging.DEBUG)
|
| 906 |
+
monkeypatch.setattr(base_rt_mod, "_RESPONSE_REJECTION_RETRY_DELAY", 0.01)
|
| 907 |
|
| 908 |
FakeCCE = type("FakeCCE", (Exception,), {})
|
| 909 |
+
monkeypatch.setattr(base_rt_mod, "ConnectionClosedError", FakeCCE)
|
| 910 |
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 911 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda default=OPENAI_DEFAULT_VOICE: "alloy")
|
| 912 |
+
monkeypatch.setattr(rt_mod, "get_active_tool_specs", lambda _: [])
|
| 913 |
|
| 914 |
N_TOOL_RESULTS = 400
|
| 915 |
REJECT_CALL_NUMBERS = {1, 3, 5, 10, 25, 50, 75, 100, 150, 200, 300, 399}
|
| 916 |
EXPECTED_TOTAL_CALLS = N_TOOL_RESULTS + len(REJECT_CALL_NUMBERS)
|
| 917 |
|
|
|
|
| 918 |
response_create_log: list[tuple[int, dict[str, Any]]] = []
|
| 919 |
+
handler_ref: list[rt_mod.OpenaiRealtimeHandler] = []
|
| 920 |
|
| 921 |
# ---- Fake event / error objects mirroring the OpenAI SDK shapes ----
|
| 922 |
|
|
|
|
| 940 |
for k, v in kwargs.items():
|
| 941 |
setattr(self, k, v)
|
| 942 |
|
| 943 |
+
event_queue: asyncio.Queue[FakeEvent | None] = asyncio.Queue()
|
| 944 |
+
|
| 945 |
# ---- Fake connection components ----
|
| 946 |
|
| 947 |
class FakeResponseAPI:
|
|
|
|
| 1044 |
return self
|
| 1045 |
|
| 1046 |
async def __anext__(self) -> FakeEvent:
|
| 1047 |
+
event = await event_queue.get()
|
| 1048 |
if event is None: # sentinel → end iteration
|
| 1049 |
raise StopAsyncIteration
|
| 1050 |
return event
|
|
|
|
| 1058 |
self.realtime = FakeRealtime()
|
| 1059 |
|
| 1060 |
monkeypatch.setattr(rt_mod, "AsyncOpenAI", FakeClient)
|
| 1061 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "openai")
|
| 1062 |
|
| 1063 |
# Patch dispatch_tool_call so tools complete with a result.
|
| 1064 |
async def _fake_dispatch(tool_name: str, args_json: str, deps: Any, **_kw: Any) -> dict[str, Any]:
|
|
|
|
| 1089 |
is_idle_tool_call=False,
|
| 1090 |
)
|
| 1091 |
|
| 1092 |
+
# Wait until spawned tool tasks, the listener, and the sender have drained.
|
| 1093 |
+
# This stress test queues hundreds of serialized response.create calls; a
|
| 1094 |
+
# condition-based wait avoids racing slower CI runners while still failing
|
| 1095 |
+
# promptly if the sender stops making progress.
|
| 1096 |
+
deadline = asyncio.get_event_loop().time() + 25.0
|
| 1097 |
+
while fake_response_api._call_count < EXPECTED_TOTAL_CALLS and asyncio.get_event_loop().time() < deadline:
|
| 1098 |
+
await asyncio.sleep(0.05)
|
| 1099 |
|
| 1100 |
# ---- Tear down ----
|
| 1101 |
|
|
|
|
| 1148 |
"""
|
| 1149 |
caplog.set_level(logging.DEBUG)
|
| 1150 |
|
| 1151 |
+
monkeypatch.setattr(base_rt_mod, "_RESPONSE_DONE_TIMEOUT", 0.3)
|
| 1152 |
|
| 1153 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 1154 |
handler = rt_mod.OpenaiRealtimeHandler(deps)
|
|
|
|
| 1200 |
"""
|
| 1201 |
caplog.set_level(logging.DEBUG)
|
| 1202 |
|
| 1203 |
+
monkeypatch.setattr(base_rt_mod, "_RESPONSE_DONE_TIMEOUT", 0.3)
|
| 1204 |
|
| 1205 |
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock())
|
| 1206 |
handler = rt_mod.OpenaiRealtimeHandler(deps)
|
|
|
|
| 1236 |
|
| 1237 |
timeout_logs = [r for r in caplog.records if "Timed out waiting for previous response" in r.getMessage()]
|
| 1238 |
assert len(timeout_logs) == 1, f"Expected 1 pre-condition timeout warning, got {len(timeout_logs)}"
|
| 1239 |
+
|
| 1240 |
+
|
| 1241 |
+
@pytest.mark.asyncio
|
| 1242 |
+
async def test_openai_excludes_head_tracking_when_no_head_tracker(monkeypatch: Any) -> None:
|
| 1243 |
+
"""head_tracking tool must not appear in OpenAI session config when head_tracker is not active."""
|
| 1244 |
+
monkeypatch.setattr(rt_mod, "get_session_instructions", lambda: "test")
|
| 1245 |
+
monkeypatch.setattr(rt_mod, "get_session_voice", lambda default=None: "alloy")
|
| 1246 |
+
|
| 1247 |
+
# mock ALL_TOOL_SPECS to include at least head_tracking and one other tool, to verify that only head_tracking is excluded, not all tools
|
| 1248 |
+
monkeypatch.setattr(
|
| 1249 |
+
ct_mod,
|
| 1250 |
+
"ALL_TOOL_SPECS",
|
| 1251 |
+
[
|
| 1252 |
+
{"type": "function", "name": "head_tracking", "description": "head_tracking", "parameters": {}},
|
| 1253 |
+
{"type": "function", "name": "fake_tool", "description": "fake_tool", "parameters": {}},
|
| 1254 |
+
],
|
| 1255 |
+
)
|
| 1256 |
+
|
| 1257 |
+
session_kwargs: dict = {}
|
| 1258 |
+
|
| 1259 |
+
class FakeSession:
|
| 1260 |
+
async def update(self, **kwargs: Any) -> None:
|
| 1261 |
+
session_kwargs["session"] = kwargs.get("session")
|
| 1262 |
+
|
| 1263 |
+
class FakeInputAudioBuffer:
|
| 1264 |
+
async def append(self, **_kw: Any) -> None:
|
| 1265 |
+
pass
|
| 1266 |
+
|
| 1267 |
+
class FakeItem:
|
| 1268 |
+
async def create(self, **_kw: Any) -> None:
|
| 1269 |
+
pass
|
| 1270 |
+
|
| 1271 |
+
class FakeConversation:
|
| 1272 |
+
item = FakeItem()
|
| 1273 |
+
|
| 1274 |
+
class FakeResponse:
|
| 1275 |
+
async def create(self, **_kw: Any) -> None:
|
| 1276 |
+
pass
|
| 1277 |
+
|
| 1278 |
+
async def cancel(self, **_kw: Any) -> None:
|
| 1279 |
+
pass
|
| 1280 |
+
|
| 1281 |
+
class FakeConn:
|
| 1282 |
+
session = FakeSession()
|
| 1283 |
+
input_audio_buffer = FakeInputAudioBuffer()
|
| 1284 |
+
conversation = FakeConversation()
|
| 1285 |
+
response = FakeResponse()
|
| 1286 |
+
|
| 1287 |
+
async def __aenter__(self) -> "FakeConn":
|
| 1288 |
+
return self
|
| 1289 |
+
|
| 1290 |
+
async def __aexit__(self, *_: Any) -> bool:
|
| 1291 |
+
return False
|
| 1292 |
+
|
| 1293 |
+
async def close(self) -> None:
|
| 1294 |
+
pass
|
| 1295 |
+
|
| 1296 |
+
def __aiter__(self) -> "FakeConn":
|
| 1297 |
+
return self
|
| 1298 |
+
|
| 1299 |
+
async def __anext__(self) -> Any:
|
| 1300 |
+
raise StopAsyncIteration
|
| 1301 |
+
|
| 1302 |
+
class FakeRealtime:
|
| 1303 |
+
def connect(self, **_kw: Any) -> FakeConn:
|
| 1304 |
+
return FakeConn()
|
| 1305 |
+
|
| 1306 |
+
class FakeClient:
|
| 1307 |
+
def __init__(self) -> None:
|
| 1308 |
+
self.realtime = FakeRealtime()
|
| 1309 |
+
|
| 1310 |
+
# case 1: no camera at all, --no-camera flag passed
|
| 1311 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock(), camera_worker=None)
|
| 1312 |
+
handler = OpenaiRealtimeHandler(deps)
|
| 1313 |
+
handler.client = FakeClient()
|
| 1314 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_up", MagicMock())
|
| 1315 |
+
monkeypatch.setattr(type(handler.tool_manager), "shutdown", AsyncMock())
|
| 1316 |
+
|
| 1317 |
+
await handler._run_realtime_session()
|
| 1318 |
+
|
| 1319 |
+
session_tools = session_kwargs.get("session", {}).get("tools", [])
|
| 1320 |
+
tool_names = [t["name"] for t in session_tools]
|
| 1321 |
+
assert "head_tracking" not in tool_names, "case 1 failed: camera_worker=None"
|
| 1322 |
+
assert "fake_tool" in tool_names, "case 1 failed: a non-head-tracking tool was unexpectedly excluded"
|
| 1323 |
+
|
| 1324 |
+
# case 2: camera is running but --head-tracker flag was not passed
|
| 1325 |
+
session_kwargs.clear()
|
| 1326 |
+
camera_worker = MagicMock()
|
| 1327 |
+
camera_worker.head_tracker = None
|
| 1328 |
+
deps = ToolDependencies(reachy_mini=MagicMock(), movement_manager=MagicMock(), camera_worker=camera_worker)
|
| 1329 |
+
handler = OpenaiRealtimeHandler(deps)
|
| 1330 |
+
handler.client = FakeClient()
|
| 1331 |
+
monkeypatch.setattr(type(handler.tool_manager), "start_up", MagicMock())
|
| 1332 |
+
monkeypatch.setattr(type(handler.tool_manager), "shutdown", AsyncMock())
|
| 1333 |
+
|
| 1334 |
+
await handler._run_realtime_session()
|
| 1335 |
+
|
| 1336 |
+
session_tools = session_kwargs.get("session", {}).get("tools", [])
|
| 1337 |
+
tool_names = [t["name"] for t in session_tools]
|
| 1338 |
+
assert "head_tracking" not in tool_names, "case 2 failed: camera_worker.head_tracker=None"
|
| 1339 |
+
assert "fake_tool" in tool_names, "case 2 failed: a non-head-tracking tool was unexpectedly excluded"
|
tests/test_profile_paths.py
CHANGED
|
@@ -7,8 +7,12 @@ import pytest
|
|
| 7 |
|
| 8 |
import reachy_mini_conversation_app.config as config_mod
|
| 9 |
import reachy_mini_conversation_app.prompts as prompts_mod
|
|
|
|
| 10 |
from reachy_mini_conversation_app.config import DEFAULT_PROFILES_DIRECTORY, config
|
|
|
|
| 11 |
from reachy_mini_conversation_app.headless_personality import (
|
|
|
|
|
|
|
| 12 |
resolve_profile_dir,
|
| 13 |
read_instructions_for,
|
| 14 |
)
|
|
@@ -81,6 +85,29 @@ def test_prompts_load_from_compact_builtin_profile(monkeypatch: pytest.MonkeyPat
|
|
| 81 |
assert read_instructions_for("mad_scientist_assistant") == expected
|
| 82 |
|
| 83 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
def test_session_voice_defaults_follow_selected_backend(monkeypatch: pytest.MonkeyPatch) -> None:
|
| 85 |
"""Session voice should fall back to the active backend default."""
|
| 86 |
monkeypatch.setattr(config, "BACKEND_PROVIDER", "gemini")
|
|
@@ -90,6 +117,21 @@ def test_session_voice_defaults_follow_selected_backend(monkeypatch: pytest.Monk
|
|
| 90 |
assert prompts_mod.get_session_voice() == "Kore"
|
| 91 |
|
| 92 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 93 |
def test_packaged_profiles_win_outside_source_checkout(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
| 94 |
"""Installed builds should use packaged profiles, not an unrelated sibling folder."""
|
| 95 |
unrelated_profiles = tmp_path / "profiles"
|
|
|
|
| 7 |
|
| 8 |
import reachy_mini_conversation_app.config as config_mod
|
| 9 |
import reachy_mini_conversation_app.prompts as prompts_mod
|
| 10 |
+
import reachy_mini_conversation_app.headless_personality as headless_mod
|
| 11 |
from reachy_mini_conversation_app.config import DEFAULT_PROFILES_DIRECTORY, config
|
| 12 |
+
from reachy_mini_conversation_app.gradio_personality import PersonalityUI
|
| 13 |
from reachy_mini_conversation_app.headless_personality import (
|
| 14 |
+
DEFAULT_OPTION,
|
| 15 |
+
read_tools_for,
|
| 16 |
resolve_profile_dir,
|
| 17 |
read_instructions_for,
|
| 18 |
)
|
|
|
|
| 85 |
assert read_instructions_for("mad_scientist_assistant") == expected
|
| 86 |
|
| 87 |
|
| 88 |
+
def test_builtin_default_profile_tools_load_for_ui() -> None:
|
| 89 |
+
"""The UI should read built-in default tools from the packaged default profile."""
|
| 90 |
+
expected = (DEFAULT_PROFILES_DIRECTORY / "default" / "tools.txt").read_text(encoding="utf-8")
|
| 91 |
+
|
| 92 |
+
assert read_tools_for(DEFAULT_OPTION) == expected
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def test_gradio_personality_ui_prefills_builtin_default_tools(monkeypatch: pytest.MonkeyPatch) -> None:
|
| 96 |
+
"""Gradio should show the built-in default profile tools on first render."""
|
| 97 |
+
monkeypatch.setattr(config, "REACHY_MINI_CUSTOM_PROFILE", None)
|
| 98 |
+
|
| 99 |
+
ui = PersonalityUI()
|
| 100 |
+
ui.create_components()
|
| 101 |
+
|
| 102 |
+
expected_tools = read_tools_for(ui.DEFAULT_OPTION)
|
| 103 |
+
expected_enabled = [
|
| 104 |
+
line.strip() for line in expected_tools.splitlines() if line.strip() and not line.strip().startswith("#")
|
| 105 |
+
]
|
| 106 |
+
|
| 107 |
+
assert ui.tools_txt_ta.value == expected_tools
|
| 108 |
+
assert sorted(ui.available_tools_cg.value) == sorted(expected_enabled)
|
| 109 |
+
|
| 110 |
+
|
| 111 |
def test_session_voice_defaults_follow_selected_backend(monkeypatch: pytest.MonkeyPatch) -> None:
|
| 112 |
"""Session voice should fall back to the active backend default."""
|
| 113 |
monkeypatch.setattr(config, "BACKEND_PROVIDER", "gemini")
|
|
|
|
| 117 |
assert prompts_mod.get_session_voice() == "Kore"
|
| 118 |
|
| 119 |
|
| 120 |
+
def test_headless_profile_write_defaults_voice_at_call_time(
|
| 121 |
+
tmp_path: Path,
|
| 122 |
+
monkeypatch: pytest.MonkeyPatch,
|
| 123 |
+
) -> None:
|
| 124 |
+
"""New headless profiles should use the currently selected backend default voice."""
|
| 125 |
+
monkeypatch.setattr(config, "BACKEND_PROVIDER", "gemini")
|
| 126 |
+
monkeypatch.setattr(config, "MODEL_NAME", "gemini-3.1-flash-live-preview")
|
| 127 |
+
monkeypatch.setattr(headless_mod, "_profiles_root", lambda: tmp_path)
|
| 128 |
+
|
| 129 |
+
headless_mod._write_profile("runtime_voice_default", "test instructions", "")
|
| 130 |
+
|
| 131 |
+
voice_file = tmp_path / "user_personalities" / "runtime_voice_default" / "voice.txt"
|
| 132 |
+
assert voice_file.read_text(encoding="utf-8") == "Kore\n"
|
| 133 |
+
|
| 134 |
+
|
| 135 |
def test_packaged_profiles_win_outside_source_checkout(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
| 136 |
"""Installed builds should use packaged profiles, not an unrelated sibling folder."""
|
| 137 |
unrelated_profiles = tmp_path / "profiles"
|
tests/test_startup_settings.py
ADDED
|
@@ -0,0 +1,80 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Tests for persisted instance-local startup settings."""
|
| 2 |
+
|
| 3 |
+
from reachy_mini_conversation_app.startup_settings import (
|
| 4 |
+
StartupSettings,
|
| 5 |
+
read_startup_settings,
|
| 6 |
+
write_startup_settings,
|
| 7 |
+
load_startup_settings_into_runtime,
|
| 8 |
+
)
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def test_write_and_read_startup_settings(tmp_path) -> None:
|
| 12 |
+
"""Startup settings should round-trip through startup_settings.json."""
|
| 13 |
+
write_startup_settings(tmp_path, profile="sorry_bro", voice="shimmer")
|
| 14 |
+
|
| 15 |
+
assert read_startup_settings(tmp_path) == StartupSettings(profile="sorry_bro", voice="shimmer")
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def test_load_startup_settings_into_runtime_applies_profile_when_no_env(monkeypatch, tmp_path) -> None:
|
| 19 |
+
"""Startup settings should seed the runtime profile when no explicit env override exists."""
|
| 20 |
+
write_startup_settings(tmp_path, profile="sorry_bro", voice="shimmer")
|
| 21 |
+
applied_profiles: list[str | None] = []
|
| 22 |
+
monkeypatch.delenv("REACHY_MINI_CUSTOM_PROFILE", raising=False)
|
| 23 |
+
monkeypatch.setattr(
|
| 24 |
+
"reachy_mini_conversation_app.config.set_custom_profile",
|
| 25 |
+
lambda profile: applied_profiles.append(profile),
|
| 26 |
+
)
|
| 27 |
+
|
| 28 |
+
settings = load_startup_settings_into_runtime(tmp_path)
|
| 29 |
+
|
| 30 |
+
assert settings == StartupSettings(profile="sorry_bro", voice="shimmer")
|
| 31 |
+
assert applied_profiles == ["sorry_bro"]
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def test_load_startup_settings_into_runtime_saved_settings_override_instance_env(monkeypatch, tmp_path) -> None:
|
| 35 |
+
"""Saved startup settings should override an instance-local profile env value."""
|
| 36 |
+
write_startup_settings(tmp_path, profile="sorry_bro", voice="shimmer")
|
| 37 |
+
applied_profiles: list[str | None] = []
|
| 38 |
+
monkeypatch.setenv("REACHY_MINI_CUSTOM_PROFILE", "env_profile")
|
| 39 |
+
monkeypatch.setattr(
|
| 40 |
+
"reachy_mini_conversation_app.config.set_custom_profile",
|
| 41 |
+
lambda profile: applied_profiles.append(profile),
|
| 42 |
+
)
|
| 43 |
+
|
| 44 |
+
settings = load_startup_settings_into_runtime(tmp_path)
|
| 45 |
+
|
| 46 |
+
assert settings == StartupSettings(profile="sorry_bro", voice="shimmer")
|
| 47 |
+
assert applied_profiles == ["sorry_bro"]
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def test_load_startup_settings_into_runtime_saved_settings_override_inherited_env(monkeypatch, tmp_path) -> None:
|
| 51 |
+
"""Saved startup settings should override a profile inherited from another `.env`."""
|
| 52 |
+
write_startup_settings(tmp_path, profile="nature_documentarian", voice="cedar")
|
| 53 |
+
applied_profiles: list[str | None] = []
|
| 54 |
+
monkeypatch.setenv("REACHY_MINI_CUSTOM_PROFILE", "example")
|
| 55 |
+
monkeypatch.setattr(
|
| 56 |
+
"reachy_mini_conversation_app.config.set_custom_profile",
|
| 57 |
+
lambda profile: applied_profiles.append(profile),
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
settings = load_startup_settings_into_runtime(tmp_path)
|
| 61 |
+
|
| 62 |
+
assert settings == StartupSettings(profile="nature_documentarian", voice="cedar")
|
| 63 |
+
assert applied_profiles == ["nature_documentarian"]
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def test_load_startup_settings_into_runtime_preserves_inherited_env_without_saved_settings(
|
| 67 |
+
monkeypatch, tmp_path
|
| 68 |
+
) -> None:
|
| 69 |
+
"""Inherited env config should still apply when no startup settings have been saved."""
|
| 70 |
+
applied_profiles: list[str | None] = []
|
| 71 |
+
monkeypatch.setenv("REACHY_MINI_CUSTOM_PROFILE", "example")
|
| 72 |
+
monkeypatch.setattr(
|
| 73 |
+
"reachy_mini_conversation_app.config.set_custom_profile",
|
| 74 |
+
lambda profile: applied_profiles.append(profile),
|
| 75 |
+
)
|
| 76 |
+
|
| 77 |
+
settings = load_startup_settings_into_runtime(tmp_path)
|
| 78 |
+
|
| 79 |
+
assert settings == StartupSettings()
|
| 80 |
+
assert applied_profiles == []
|