From 369f6012b09f6105cf3b543d03fb6994c937745a Mon Sep 17 00:00:00 2001 From: Gijs de Jong Date: Wed, 23 Sep 2026 13:56:02 +0200 Subject: [PATCH 1/3] 1.8 support --- README.md | 14 +- booster_sdk/src/client/ai.rs | 491 +++++++++++++++++- booster_sdk/src/client/audio.rs | 14 + booster_sdk/src/client/loco.rs | 70 ++- booster_sdk/src/dds/mod.rs | 2 + booster_sdk/src/dds/operation.rs | 456 ++++++++++++++++ booster_sdk/src/dds/rpc.rs | 34 +- booster_sdk/src/dds/topics.rs | 3 + booster_sdk/src/types/error.rs | 4 + booster_sdk_py/booster_sdk/client/lui.py | 8 + .../booster_sdk_bindings.pyi | 155 +++++- booster_sdk_py/src/client/ai.rs | 8 + booster_sdk_py/src/client/audio.rs | 5 +- booster_sdk_py/src/client/booster.rs | 9 + booster_sdk_py/src/client/lui.rs | 203 +++++++- 15 files changed, 1441 insertions(+), 35 deletions(-) create mode 100644 booster_sdk/src/dds/operation.rs diff --git a/README.md b/README.md index 1064659..33d19cf 100644 --- a/README.md +++ b/README.md @@ -37,11 +37,23 @@ client.move_robot(0.5, 0.0, 0.0) ``` -The Rust crate and Python bindings support the Booster 1.7 SDK, including +The Rust crate and Python bindings support the Booster 1.8 SDK, including locomotion, robot/device catalog discovery, gripper and LED control, AI/LUI and vision RPCs, camera discovery, hand-eye calibration, and audio device/Bluetooth management. +SDK 1.8 additions include: + +- asynchronous RPC operations with progress/result events and cancellation + (`BoosterClient::start_operation` / `cancel_operation`, Rust only); +- resetting odometry to a target pose (`reset_odometry_to`); +- per-client LUI ASR/TTS sessions, speech synthesis to audio, and one-shot or + session-based audio recognition; +- an optional AgentHub `persona_id` for AI chat; +- the extended LUI ASR chunk message (session, timing, speaker and utterance + details); +- the low-battery RPC status (503) and a 3-channel default raw capture format. + SDK 1.7 additions include: - timed head rotation and selectable v1/v2 get-up behavior; diff --git a/booster_sdk/src/client/ai.rs b/booster_sdk/src/client/ai.rs index 62d42a0..909a07e 100644 --- a/booster_sdk/src/client/ai.rs +++ b/booster_sdk/src/client/ai.rs @@ -1,14 +1,16 @@ //! AI and LUI high-level RPC clients. +use std::sync::{Arc, Mutex}; use std::time::Duration; use serde::{Deserialize, Serialize}; +use uuid::Uuid; use crate::dds::{ - AI_API_TOPIC, DdsNode, DdsSubscription, LUI_API_TOPIC, RpcClient, RpcClientOptions, - ai_subtitle_topic, lui_asr_chunk_topic, + AI_API_TOPIC, DdsNode, DdsSubscription, LUI_API_OPERATION_EVENT_TOPIC, LUI_API_TOPIC, + RpcClient, RpcClientOptions, RpcOperationClient, ai_subtitle_topic, lui_asr_chunk_topic, }; -use crate::types::Result; +use crate::types::{BoosterError, Result, RpcError}; crate::api_id_enum! { /// AI chat RPC API identifiers. @@ -29,9 +31,25 @@ crate::api_id_enum! { StartTts = 1050, StopTts = 1051, SendTtsText = 1052, + SynthesizeSpeech = 1100, + RecognizeAudioOnce = 1101, + StartAudioRecognizer = 1102, + StopAudioRecognizer = 1103, + RecognizeAudioInSession = 1104, } } +/// Maximum text length for [`LuiClient::synthesize_speech`], in Unicode code points. +pub const LUI_MAX_SYNTHESIZE_TEXT_CODE_POINTS: usize = 1000; +/// Maximum audio duration for LUI recognition requests, in milliseconds. +pub const LUI_MAX_RECOGNIZE_AUDIO_DURATION_MS: u64 = 120 * 1000; +/// Maximum decoded audio size for LUI recognition requests, in bytes. +pub const LUI_MAX_RECOGNIZE_AUDIO_BYTES: u64 = 6 * 1024 * 1024; + +const LUI_OPERATION_START_TIMEOUT: Duration = Duration::from_millis(3000); +const LUI_SYNTHESIZE_TIMEOUT: Duration = Duration::from_secs(1800); +const LUI_RECOGNIZE_TIMEOUT: Duration = Duration::from_secs(900); + /// TTS configuration for AI chat. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct TtsConfig { @@ -57,6 +75,9 @@ pub struct AsrConfig { /// Parameters for starting AI chat. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct StartAiChatParameter { + /// Optional AgentHub persona identifier. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub persona_id: Option, pub interrupt_mode: bool, pub asr_config: AsrConfig, pub llm_config: LlmConfig, @@ -95,10 +116,123 @@ pub struct Subtitle { pub round_id: i32, } -/// LUI ASR chunk topic payload. +/// Speech synthesis request for [`LuiClient::synthesize_speech`]. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct LuiSynthesizeSpeechRequest { + pub text: String, + pub voice_type: String, + /// Speech speed ratio: 1.0 = normal, 0.5 = slowest, 2.0 = fastest. + pub speed: f64, + /// Whether the robot should also play the synthesized audio. + pub playback: bool, +} + +impl LuiSynthesizeSpeechRequest { + /// Request with the default voice, normal speed and no playback. + pub fn new(text: impl Into) -> Self { + Self { + text: text.into(), + voice_type: "default".to_owned(), + speed: 1.0, + playback: false, + } + } +} + +/// Synthesized audio returned by [`LuiClient::synthesize_speech`]. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct LuiSynthesizeSpeechResponse { + pub audio_base64: String, + pub sample_rate_hz: i32, + pub channels: i32, + pub bits_per_sample: i32, + pub format: String, +} + +/// Audio recognition request for the LUI recognize APIs. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct LuiRecognizeAudioRequest { + /// `"file"` or `"pcm"`. + pub input_type: String, + #[serde(default, skip_serializing_if = "String::is_empty")] + pub file_path: String, + #[serde(default, skip_serializing_if = "String::is_empty")] + pub audio_base64: String, + pub sample_rate_hz: i32, + pub channels: i32, + pub bits_per_sample: i32, + pub format: String, +} + +impl LuiRecognizeAudioRequest { + /// Recognize a `.wav` or `.mp3` file on the robot. + pub fn from_file(file_path: impl Into) -> Self { + Self { + input_type: "file".to_owned(), + file_path: file_path.into(), + audio_base64: String::new(), + sample_rate_hz: 16000, + channels: 1, + bits_per_sample: 16, + format: "pcm_s16le".to_owned(), + } + } + + /// Recognize base64-encoded 16 kHz mono 16-bit little-endian raw PCM. + pub fn from_pcm_base64(audio_base64: impl Into) -> Self { + Self { + input_type: "pcm".to_owned(), + file_path: String::new(), + audio_base64: audio_base64.into(), + sample_rate_hz: 16000, + channels: 1, + bits_per_sample: 16, + format: "pcm_s16le".to_owned(), + } + } +} + +/// Recognition result. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct LuiRecognizeAudioResponse { + pub text: String, +} + +/// Utterance entry inside an [`AsrChunk`]. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct AsrUtterance { + pub text: String, + pub definite: bool, + pub start_time_ms: u32, + pub end_time_ms: u32, + pub speaker_id: String, + pub emotion: String, + pub gender: String, + pub lid_lang: String, + pub speech_rate: f32, + pub volume_db: f32, + pub additions_json: String, +} + +/// LUI ASR chunk topic payload. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct AsrChunk { + pub session_id: String, + pub client_id: String, + pub role: String, pub text: String, + pub is_final: bool, + pub start_time_ms: u32, + pub end_time_ms: u32, + pub speaker_id: String, + pub emotion: String, + pub gender: String, + pub lid_lang: String, + pub speech_rate: f32, + pub volume_db: f32, + pub provider_logid: String, + pub raw_json: String, + pub utterances: Vec, } /// User identifier used by robot-generated subtitle entries. @@ -164,9 +298,22 @@ impl AiClient { } } +#[derive(Debug, Default)] +struct LuiSessions { + asr: Option, + audio_recognizer: Option, + tts: Option<(String, String)>, +} + /// High-level RPC client for LUI ASR/TTS features. +/// +/// Since SDK 1.8 every ASR, TTS and recognizer session is owned by the client +/// instance that started it, identified by a per-client id. pub struct LuiClient { rpc: RpcClient, + client_id: String, + sessions: Mutex, + operations: Mutex>>, } impl LuiClient { @@ -185,7 +332,12 @@ impl LuiClient { /// Create a LUI client with custom RPC options. pub fn with_options(options: RpcClientOptions) -> Result { let rpc = RpcClient::for_topic(options, LUI_API_TOPIC)?; - Ok(Self { rpc }) + Ok(Self { + rpc, + client_id: Uuid::new_v4().to_string(), + sessions: Mutex::new(LuiSessions::default()), + operations: Mutex::new(None), + }) } /// Access the underlying DDS node. @@ -193,29 +345,250 @@ impl LuiClient { self.rpc.node() } - /// Start ASR. + /// Identifier of this client instance, sent with every session request. + pub fn client_id(&self) -> &str { + &self.client_id + } + + /// Session id of the ASR session started by this client, if any. + pub fn current_asr_session_id(&self) -> Option { + self.sessions().ok()?.asr.clone() + } + + /// Session id of the audio recognizer started by this client, if any. + pub fn current_audio_recognizer_session_id(&self) -> Option { + self.sessions().ok()?.audio_recognizer.clone() + } + + /// Session id of the TTS session started by this client, if any. + pub fn current_tts_session_id(&self) -> Option { + self.sessions().ok()?.tts.as_ref().map(|(id, _)| id.clone()) + } + + fn sessions(&self) -> Result> { + self.sessions + .lock() + .map_err(|_| BoosterError::Other("LUI session state poisoned".to_owned())) + } + + fn session_body(&self, session_id: &str) -> serde_json::Value { + serde_json::json!({ "session_id": session_id, "client_id": self.client_id }) + } + + fn operation_client(&self) -> Result> { + let mut operations = self + .operations + .lock() + .map_err(|_| BoosterError::Other("operation client poisoned".to_owned()))?; + if let Some(client) = operations.as_ref() { + return Ok(Arc::clone(client)); + } + let client = Arc::new(RpcOperationClient::new( + self.rpc.node(), + LUI_API_OPERATION_EVENT_TOPIC, + )?); + *operations = Some(Arc::clone(&client)); + Ok(client) + } + + async fn run_operation(&self, api_id: LuiApiId, body: String, timeout: Duration) -> Result + where + R: serde::de::DeserializeOwned, + { + let mut operation = self + .operation_client()? + .start(&self.rpc, api_id.into(), body, LUI_OPERATION_START_TIMEOUT) + .await?; + let body = operation.wait_timeout(timeout).await?.into_body()?; + Ok(serde_json::from_str(body.trim())?) + } + + /// Start ASR. Recognized text is published on the ASR chunk topic. + /// + /// Does nothing if this client already started ASR. pub async fn start_asr(&self) -> Result<()> { - self.rpc.call_void(LuiApiId::StartAsr, "").await + if self.sessions()?.asr.is_some() { + return Ok(()); + } + let session_id = Uuid::new_v4().to_string(); + self.rpc + .call_void( + LuiApiId::StartAsr, + self.session_body(&session_id).to_string(), + ) + .await?; + self.sessions()?.asr = Some(session_id); + Ok(()) } - /// Stop ASR. + /// Stop the ASR session started by this client. pub async fn stop_asr(&self) -> Result<()> { - self.rpc.call_void(LuiApiId::StopAsr, "").await + let Some(session_id) = self.sessions()?.asr.clone() else { + return Ok(()); + }; + self.rpc + .call_void( + LuiApiId::StopAsr, + self.session_body(&session_id).to_string(), + ) + .await?; + self.sessions()?.asr = None; + Ok(()) } - /// Start TTS with the given configuration. + /// Start a TTS session owned by this client. + /// + /// Succeeds without a request if this client already has a TTS session + /// with the same voice; fails with a conflict for a different voice. pub async fn start_tts(&self, config: &LuiTtsConfig) -> Result<()> { - self.rpc.call_serialized(LuiApiId::StartTts, config).await + if let Some((_, voice_type)) = &self.sessions()?.tts { + if *voice_type == config.voice_type { + return Ok(()); + } + return Err(RpcError::Conflict(format!( + "TTS already started with voice type '{voice_type}'" + )) + .into()); + } + let session_id = Uuid::new_v4().to_string(); + let mut body = self.session_body(&session_id); + body["voice_type"] = config.voice_type.clone().into(); + self.rpc + .call_void(LuiApiId::StartTts, body.to_string()) + .await?; + self.sessions()?.tts = Some((session_id, config.voice_type.clone())); + Ok(()) } - /// Stop TTS. + /// Stop the TTS session started by this client. pub async fn stop_tts(&self) -> Result<()> { - self.rpc.call_void(LuiApiId::StopTts, "").await + let Some((session_id, _)) = self.sessions()?.tts.clone() else { + return Ok(()); + }; + self.rpc + .call_void( + LuiApiId::StopTts, + self.session_body(&session_id).to_string(), + ) + .await?; + self.sessions()?.tts = None; + Ok(()) + } + + fn tts_session_id(&self) -> Result { + self.sessions()? + .tts + .as_ref() + .map(|(id, _)| id.clone()) + .ok_or_else(|| RpcError::BadRequest("TTS is not started".to_owned()).into()) } - /// Send text to TTS. + /// Queue text in this client's TTS session. + /// + /// Requires [`Self::start_tts`]. The service rejects requests sent more + /// than once per second. pub async fn send_tts_text(&self, param: &LuiTtsParameter) -> Result<()> { - self.rpc.call_serialized(LuiApiId::SendTtsText, param).await + let mut body = self.session_body(&self.tts_session_id()?); + body["text"] = param.text.clone().into(); + self.rpc + .call_void(LuiApiId::SendTtsText, body.to_string()) + .await + } + + /// Synthesize speech in this client's TTS session and return the audio. + /// + /// Requires [`Self::start_tts`]. + pub async fn synthesize_speech( + &self, + req: &LuiSynthesizeSpeechRequest, + ) -> Result { + let text_len = req.text.chars().count(); + if text_len > LUI_MAX_SYNTHESIZE_TEXT_CODE_POINTS { + return Err(RpcError::BadRequest(format!( + "text length {text_len} exceeds limit {LUI_MAX_SYNTHESIZE_TEXT_CODE_POINTS}" + )) + .into()); + } + let session_id = self.tts_session_id()?; + let mut body = serde_json::to_value(req)?; + body["session_id"] = session_id.into(); + body["client_id"] = self.client_id.clone().into(); + self.run_operation( + LuiApiId::SynthesizeSpeech, + body.to_string(), + LUI_SYNTHESIZE_TIMEOUT, + ) + .await + } + + /// Recognize an audio file or PCM payload without a session. + pub async fn recognize_audio_once( + &self, + req: &LuiRecognizeAudioRequest, + ) -> Result { + validate_recognize_audio(req)?; + self.run_operation( + LuiApiId::RecognizeAudioOnce, + serde_json::to_string(req)?, + LUI_RECOGNIZE_TIMEOUT, + ) + .await + } + + /// Start a reusable audio recognizer session owned by this client. + /// + /// Does nothing if this client already started one. + pub async fn start_audio_recognizer(&self) -> Result<()> { + if self.sessions()?.audio_recognizer.is_some() { + return Ok(()); + } + let session_id = Uuid::new_v4().to_string(); + self.rpc + .call_void( + LuiApiId::StartAudioRecognizer, + self.session_body(&session_id).to_string(), + ) + .await?; + self.sessions()?.audio_recognizer = Some(session_id); + Ok(()) + } + + /// Stop the audio recognizer session started by this client. + pub async fn stop_audio_recognizer(&self) -> Result<()> { + let Some(session_id) = self.sessions()?.audio_recognizer.clone() else { + return Ok(()); + }; + self.rpc + .call_void( + LuiApiId::StopAudioRecognizer, + self.session_body(&session_id).to_string(), + ) + .await?; + self.sessions()?.audio_recognizer = None; + Ok(()) + } + + /// Recognize audio in this client's audio recognizer session. + /// + /// Requires [`Self::start_audio_recognizer`]. + pub async fn recognize_audio_in_session( + &self, + req: &LuiRecognizeAudioRequest, + ) -> Result { + validate_recognize_audio(req)?; + let session_id = + self.sessions()?.audio_recognizer.clone().ok_or_else(|| { + RpcError::BadRequest("audio recognizer is not started".to_owned()) + })?; + let mut body = serde_json::to_value(req)?; + body["session_id"] = session_id.into(); + body["client_id"] = self.client_id.clone().into(); + self.run_operation( + LuiApiId::RecognizeAudioInSession, + body.to_string(), + LUI_RECOGNIZE_TIMEOUT, + ) + .await } /// Subscribe to ASR chunk messages. @@ -223,3 +596,91 @@ impl LuiClient { self.rpc.node().subscribe(&lui_asr_chunk_topic(), 16) } } + +/// Decoded byte count of a base64 payload, derived from its length. +fn estimate_base64_decoded_bytes(base64: &str) -> u64 { + let padding = base64.bytes().rev().take_while(|&b| b == b'=').count() as u64; + (base64.len() as u64 / 4 * 3).saturating_sub(padding) +} + +/// Client-side check of the PCM limits the LUI service enforces. +fn validate_recognize_audio(req: &LuiRecognizeAudioRequest) -> Result<()> { + if req.input_type != "pcm" { + return Ok(()); + } + let bytes = estimate_base64_decoded_bytes(&req.audio_base64); + if bytes > LUI_MAX_RECOGNIZE_AUDIO_BYTES { + return Err(RpcError::BadRequest(format!( + "audio size {bytes} bytes exceeds limit {LUI_MAX_RECOGNIZE_AUDIO_BYTES} bytes" + )) + .into()); + } + let bytes_per_second = u64::try_from( + i64::from(req.sample_rate_hz) * i64::from(req.channels) * i64::from(req.bits_per_sample) + / 8, + ) + .unwrap_or(0); + if let Some(duration_ms) = (bytes * 1000).checked_div(bytes_per_second) + && duration_ms > LUI_MAX_RECOGNIZE_AUDIO_DURATION_MS + { + return Err(RpcError::BadRequest(format!( + "audio duration {duration_ms} ms exceeds limit {LUI_MAX_RECOGNIZE_AUDIO_DURATION_MS} ms" + )) + .into()); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn base64_decoded_size_accounts_for_padding() { + assert_eq!(estimate_base64_decoded_bytes("QUJD"), 3); + assert_eq!(estimate_base64_decoded_bytes("QUI="), 2); + assert_eq!(estimate_base64_decoded_bytes("QQ=="), 1); + } + + #[test] + fn recognize_request_omits_empty_fields() { + let json = serde_json::to_value(LuiRecognizeAudioRequest::from_file("/tmp/a.wav")).unwrap(); + assert_eq!(json["input_type"], "file"); + assert_eq!(json["file_path"], "/tmp/a.wav"); + assert!(json.get("audio_base64").is_none()); + } + + #[test] + fn rejects_too_long_pcm() { + // 16 kHz mono s16 = 32000 B/s; 121 s of audio. + let bytes = 32000 * 121; + let req = LuiRecognizeAudioRequest::from_pcm_base64("A".repeat(bytes / 3 * 4)); + assert!(validate_recognize_audio(&req).is_err()); + let req = LuiRecognizeAudioRequest::from_pcm_base64("AAAA"); + assert!(validate_recognize_audio(&req).is_ok()); + } + + #[test] + fn start_ai_chat_omits_missing_persona() { + let param = StartAiChatParameter { + persona_id: None, + interrupt_mode: false, + asr_config: AsrConfig { + interrupt_speech_duration: 0, + interrupt_keywords: vec![], + }, + llm_config: LlmConfig { + system_prompt: String::new(), + welcome_msg: String::new(), + prompt_name: String::new(), + }, + tts_config: TtsConfig { + voice_type: String::new(), + ignore_bracket_text: vec![], + }, + enable_face_tracking: false, + }; + let json = serde_json::to_value(¶m).unwrap(); + assert!(json.get("persona_id").is_none()); + } +} diff --git a/booster_sdk/src/client/audio.rs b/booster_sdk/src/client/audio.rs index 0c2bac8..063c3eb 100644 --- a/booster_sdk/src/client/audio.rs +++ b/booster_sdk/src/client/audio.rs @@ -294,6 +294,20 @@ pub struct AudioCaptureStreamOptions { pub requested_raw_format: PcmFormat, } +impl Default for AudioCaptureStreamOptions { + fn default() -> Self { + Self { + enable_raw_pcm: true, + enable_naec_pcm: false, + // BoosterAEC raw PCM default: 16 kHz, 3 channels, 16-bit. + requested_raw_format: PcmFormat { + channels: 3, + ..PcmFormat::default() + }, + } + } +} + /// Generic audio service result. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct ServiceResult { diff --git a/booster_sdk/src/client/loco.rs b/booster_sdk/src/client/loco.rs index c46d982..9f40928 100644 --- a/booster_sdk/src/client/loco.rs +++ b/booster_sdk/src/client/loco.rs @@ -1,14 +1,16 @@ //! High-level B1 locomotion client built on DDS RPC and topic I/O. +use std::sync::{Arc, Mutex}; use std::time::Duration; use crate::dds::{ - BatteryState, BinaryData, ButtonEventMsg, DdsNode, DdsPublisher, DdsSubscription, - GripperControl, LightControlMsg, MotionState, RemoteControllerState, RobotProcessStateMsg, - RobotStatusDdsMsg, RpcClient, RpcClientOptions, SafeMode, battery_state_topic, - button_event_topic, device_gateway_topic, gripper_control_topic, light_control_topic, - motion_state_topic, process_state_topic, remote_controller_topic, safe_mode_topic, - video_stream_topic, + ApiOperation, ApiOperationHandle, BatteryState, BinaryData, ButtonEventMsg, + DEFAULT_OPERATION_START_TIMEOUT, DdsNode, DdsPublisher, DdsSubscription, GripperControl, + LOCO_API_OPERATION_EVENT_TOPIC, LightControlMsg, MotionState, RemoteControllerState, + RobotProcessStateMsg, RobotStatusDdsMsg, RpcClient, RpcClientOptions, RpcOperationClient, + SafeMode, battery_state_topic, button_event_topic, device_gateway_topic, gripper_control_topic, + light_control_topic, motion_state_topic, process_state_topic, remote_controller_topic, + safe_mode_topic, video_stream_topic, }; use crate::types::{ BoosterHandType, CustomTrainedTraj, DanceId, DeviceInfo, DeviceInfoKind, @@ -24,6 +26,7 @@ use typed_builder::TypedBuilder; /// High-level client for B1 locomotion control and telemetry. pub struct BoosterClient { rpc: RpcClient, + operations: Mutex>>, gripper_publisher: DdsPublisher, light_publisher: DdsPublisher, safe_mode_publisher: DdsPublisher, @@ -50,6 +53,7 @@ impl BoosterClient { Ok(Self { rpc, + operations: Mutex::new(None), gripper_publisher, light_publisher, safe_mode_publisher, @@ -61,6 +65,54 @@ impl BoosterClient { self.rpc.node() } + fn operation_client(&self) -> Result> { + let mut operations = self + .operations + .lock() + .map_err(|_| crate::types::BoosterError::Other("operation client poisoned".into()))?; + if let Some(client) = operations.as_ref() { + return Ok(Arc::clone(client)); + } + let client = Arc::new(RpcOperationClient::new( + self.rpc.node(), + LOCO_API_OPERATION_EVENT_TOPIC, + )?); + *operations = Some(Arc::clone(&client)); + Ok(client) + } + + /// Start a locomotion API request as an asynchronous operation. + /// + /// Returns once the service has accepted the operation. Progress and the + /// final result are delivered through the returned [`ApiOperation`]. + pub async fn start_operation( + &self, + api_id: LocoApiId, + param: impl Into, + ) -> Result { + self.start_operation_with_timeout(api_id, param, DEFAULT_OPERATION_START_TIMEOUT) + .await + } + + /// Like [`Self::start_operation`], waiting up to `start_timeout` for acceptance. + pub async fn start_operation_with_timeout( + &self, + api_id: LocoApiId, + param: impl Into, + start_timeout: Duration, + ) -> Result { + self.operation_client()? + .start(&self.rpc, api_id.into(), param, start_timeout) + .await + } + + /// Cancel an operation started with [`Self::start_operation`]. + pub async fn cancel_operation(&self, handle: &ApiOperationHandle) -> Result<()> { + self.operation_client()? + .cancel(&self.rpc, handle, None) + .await + } + /// Change the robot mode. pub async fn change_mode(&self, mode: RobotMode) -> Result<()> { let param = json!({ "mode": i32::from(mode) }).to_string(); @@ -384,6 +436,12 @@ impl BoosterClient { self.rpc.call_void(LocoApiId::ResetOdometry, "").await } + /// Reset odometry to a target pose (meters, radians). + pub async fn reset_odometry_to(&self, x: f64, y: f64, theta: f64) -> Result<()> { + let param = json!({ "x": x, "y": y, "theta": theta }).to_string(); + self.rpc.call_void(LocoApiId::ResetOdometry, param).await + } + /// Load a custom trained trajectory. pub async fn load_custom_trained_traj( &self, diff --git a/booster_sdk/src/dds/mod.rs b/booster_sdk/src/dds/mod.rs index 10b55f3..0dbcebb 100644 --- a/booster_sdk/src/dds/mod.rs +++ b/booster_sdk/src/dds/mod.rs @@ -2,11 +2,13 @@ pub mod messages; pub mod node; +pub mod operation; pub mod qos; pub mod rpc; pub mod topics; pub use messages::*; pub use node::*; +pub use operation::*; pub use rpc::*; pub use topics::*; diff --git a/booster_sdk/src/dds/operation.rs b/booster_sdk/src/dds/operation.rs new file mode 100644 index 0000000..8568b00 --- /dev/null +++ b/booster_sdk/src/dds/operation.rs @@ -0,0 +1,456 @@ +//! Asynchronous RPC operations (SDK 1.8). +//! +//! An operation is started with a regular RPC request whose header carries +//! `call_mode = 1` (operation), `operation_command = 1` (start) and a +//! client-generated `operation_id`. The start response only says whether the +//! service accepted the operation. Progress and the final result are published +//! afterwards on a separate `*OperationEvent` topic as `RpcRespMsg` samples +//! whose `uuid` is the operation id. + +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use serde_json::Value; +use tokio::sync::mpsc; +use uuid::Uuid; + +use crate::types::{Result, RpcError}; + +use super::messages::RpcRespMsg; +use super::node::DdsNode; +use super::qos::qos_reliable_keep_last; +use super::rpc::RpcClient; +use super::topics::{TYPE_RPC_RESP, TopicSpec}; + +/// Default time to wait for the service to accept an operation. +pub const DEFAULT_OPERATION_START_TIMEOUT: Duration = Duration::from_millis(1000); + +const CALL_MODE_OPERATION: i32 = 1; +const OPERATION_COMMAND_START: i32 = 1; +const OPERATION_COMMAND_CANCEL: i32 = 2; + +const EVENT_TYPE_PROGRESS: i64 = 1; +const EVENT_TYPE_FINISHED: i64 = 2; + +/// Event topic for an operation-capable RPC service. +pub fn operation_event_topic(name: &str) -> TopicSpec { + TopicSpec { + name: name.to_owned(), + type_name: TYPE_RPC_RESP, + qos: qos_reliable_keep_last(64), + kind: rustdds::TopicKind::NoKey, + } +} + +/// Final outcome of an asynchronous operation. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum ApiOperationResultCode { + Unknown, + Succeeded, + Aborted, + Canceled, + Rejected, + Timeout, +} + +impl From for ApiOperationResultCode { + fn from(value: i64) -> Self { + match value { + 1 => Self::Succeeded, + 2 => Self::Aborted, + 3 => Self::Canceled, + 4 => Self::Rejected, + 5 => Self::Timeout, + _ => Self::Unknown, + } + } +} + +impl From for i32 { + fn from(value: ApiOperationResultCode) -> Self { + match value { + ApiOperationResultCode::Unknown => 0, + ApiOperationResultCode::Succeeded => 1, + ApiOperationResultCode::Aborted => 2, + ApiOperationResultCode::Canceled => 3, + ApiOperationResultCode::Rejected => 4, + ApiOperationResultCode::Timeout => 5, + } + } +} + +/// Identifies a started operation, used to cancel it. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct ApiOperationHandle { + pub operation_id: String, + pub api_id: i32, +} + +/// Intermediate progress published while an operation runs. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ApiOperationProgress { + pub operation_id: String, + pub api_id: i64, + pub status: i64, + pub message: String, + pub body: String, +} + +/// Final result of an operation. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ApiOperationResult { + pub code: ApiOperationResultCode, + pub operation_id: String, + pub api_id: i64, + pub status: i64, + pub message: String, + pub body: String, +} + +impl ApiOperationResult { + /// Whether the operation finished successfully. + #[must_use] + pub fn is_success(&self) -> bool { + self.code == ApiOperationResultCode::Succeeded && self.status == 0 + } + + /// Convert into the result body, or an error for unsuccessful operations. + pub fn into_body(self) -> Result { + if self.is_success() { + return Ok(self.body); + } + let status = i32::try_from(self.status).unwrap_or(i32::MAX); + let message = if self.message.is_empty() { + format!("operation finished with {:?}", self.code) + } else { + self.message + }; + if status == 0 { + return Err(RpcError::RequestFailed { status, message }.into()); + } + Err(RpcError::from_status_code(status, message).into()) + } +} + +/// Event published for an operation. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ApiOperationEvent { + Progress(ApiOperationProgress), + Finished(ApiOperationResult), +} + +type Registry = Mutex>>; + +/// A started operation. Receives progress and the final result. +/// +/// Dropping it stops event delivery but does not cancel the operation on the +/// robot; use the owning client's cancel method for that. +pub struct ApiOperation { + handle: ApiOperationHandle, + events: mpsc::UnboundedReceiver, + registry: Arc, +} + +impl ApiOperation { + /// Handle identifying this operation. + pub fn handle(&self) -> &ApiOperationHandle { + &self.handle + } + + /// Operation id generated for this operation. + pub fn operation_id(&self) -> &str { + &self.handle.operation_id + } + + /// Receive the next progress or result event. + /// + /// Returns `None` after the final result has been delivered. + pub async fn next_event(&mut self) -> Option { + self.events.recv().await + } + + /// Wait for the final result, skipping progress events. + pub async fn wait(&mut self) -> Result { + while let Some(event) = self.events.recv().await { + if let ApiOperationEvent::Finished(result) = event { + return Ok(result); + } + } + Err(RpcError::RequestFailed { + status: -1, + message: "operation event stream closed".to_owned(), + } + .into()) + } + + /// Wait for the final result for at most `timeout`. + pub async fn wait_timeout(&mut self, timeout: Duration) -> Result { + tokio::time::timeout(timeout, self.wait()) + .await + .map_err(|_| RpcError::Timeout { timeout })? + } +} + +impl Drop for ApiOperation { + fn drop(&mut self) { + if let Ok(mut ops) = self.registry.lock() { + ops.remove(&self.handle.operation_id); + } + } +} + +/// Starts and cancels asynchronous operations for one RPC service and routes +/// the service's operation events to the matching [`ApiOperation`]. +pub struct RpcOperationClient { + registry: Arc, +} + +impl RpcOperationClient { + /// Subscribe to `event_topic` on `node`. + pub fn new(node: &DdsNode, event_topic: &str) -> Result { + let mut reader = + node.subscribe_reader::(&operation_event_topic(event_topic))?; + let registry: Arc = Arc::new(Mutex::new(HashMap::new())); + let weak = Arc::downgrade(®istry); + let topic = event_topic.to_owned(); + + std::thread::spawn(move || { + loop { + let Some(registry) = weak.upgrade() else { + break; + }; + match reader.take_next_sample() { + Ok(Some(sample)) => dispatch_event(®istry, &topic, sample.into_value()), + Ok(None) => { + drop(registry); + std::thread::sleep(Duration::from_millis(5)); + } + Err(_) => { + drop(registry); + std::thread::sleep(Duration::from_millis(10)); + } + } + } + }); + + Ok(Self { registry }) + } + + /// Start an operation and wait up to `start_timeout` for the service to accept it. + pub async fn start( + &self, + rpc: &RpcClient, + api_id: i32, + body: impl Into, + start_timeout: Duration, + ) -> Result { + let operation_id = Uuid::new_v4().to_string(); + let (sender, events) = mpsc::unbounded_channel(); + + // Register before sending so early events are not lost. + self.registry + .lock() + .map_err(|_| RpcError::BadRequest("operation registry poisoned".to_owned()))? + .insert(operation_id.clone(), sender); + + let operation = ApiOperation { + handle: ApiOperationHandle { + operation_id: operation_id.clone(), + api_id, + }, + events, + registry: Arc::clone(&self.registry), + }; + + let header = operation_header(api_id, OPERATION_COMMAND_START, &operation_id); + // On failure `operation` is dropped, which unregisters it. + rpc.call_with_header(api_id, header, body.into(), Some(start_timeout)) + .await?; + + Ok(operation) + } + + /// Ask the service to cancel a running operation. + /// + /// The final `Canceled` result, if any, is still delivered to the operation. + pub async fn cancel( + &self, + rpc: &RpcClient, + handle: &ApiOperationHandle, + timeout: Option, + ) -> Result<()> { + let header = operation_header( + handle.api_id, + OPERATION_COMMAND_CANCEL, + &handle.operation_id, + ); + rpc.call_with_header(handle.api_id, header, String::new(), timeout) + .await?; + Ok(()) + } +} + +fn operation_header(api_id: i32, command: i32, operation_id: &str) -> String { + serde_json::json!({ + "api_id": api_id, + "call_mode": CALL_MODE_OPERATION, + "operation_command": command, + "operation_id": operation_id, + }) + .to_string() +} + +fn parse_i64(value: Option<&Value>) -> Option { + match value? { + Value::Number(n) => n.as_i64(), + Value::String(s) => s.parse().ok(), + _ => None, + } +} + +fn parse_event(msg: RpcRespMsg) -> Option { + let header: Value = serde_json::from_str(msg.header.trim()).ok()?; + let api_id = parse_i64(header.get("api_id")).unwrap_or(0); + let status = parse_i64(header.get("status")).unwrap_or(0); + let message = header + .get("message") + .and_then(Value::as_str) + .unwrap_or_default() + .to_owned(); + + match parse_i64(header.get("event_type"))? { + EVENT_TYPE_PROGRESS => Some(ApiOperationEvent::Progress(ApiOperationProgress { + operation_id: msg.uuid, + api_id, + status, + message, + body: msg.body, + })), + EVENT_TYPE_FINISHED => Some(ApiOperationEvent::Finished(ApiOperationResult { + code: parse_i64(header.get("result_code")).map_or( + ApiOperationResultCode::Unknown, + ApiOperationResultCode::from, + ), + operation_id: msg.uuid, + api_id, + status, + message, + body: msg.body, + })), + _ => None, + } +} + +fn dispatch_event(registry: &Registry, topic: &str, msg: RpcRespMsg) { + let Ok(mut ops) = registry.lock() else { + return; + }; + let operation_id = msg.uuid.clone(); + let Some(sender) = ops.get(&operation_id) else { + return; + }; + + let Some(event) = parse_event(msg) else { + tracing::warn!( + target: "booster_sdk::rpc", + topic, + operation_id = %operation_id, + "failed to parse rpc operation event" + ); + return; + }; + + let finished = matches!(event, ApiOperationEvent::Finished(_)); + let _ = sender.send(event); + if finished { + // Dropping the sender closes the stream after the result. + ops.remove(&operation_id); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn msg(header: &str, body: &str) -> RpcRespMsg { + RpcRespMsg { + uuid: "op-1".to_owned(), + header: header.to_owned(), + body: body.to_owned(), + } + } + + #[test] + fn operation_header_contains_operation_fields() { + let header: Value = + serde_json::from_str(&operation_header(2031, OPERATION_COMMAND_START, "abc")).unwrap(); + assert_eq!(header["api_id"], 2031); + assert_eq!(header["call_mode"], 1); + assert_eq!(header["operation_command"], 1); + assert_eq!(header["operation_id"], "abc"); + } + + #[test] + fn parses_progress_event() { + let event = parse_event(msg( + r#"{"event_type":1,"api_id":1100,"status":0,"message":"working"}"#, + "{}", + )) + .unwrap(); + let ApiOperationEvent::Progress(progress) = event else { + panic!("expected progress"); + }; + assert_eq!(progress.operation_id, "op-1"); + assert_eq!(progress.api_id, 1100); + assert_eq!(progress.message, "working"); + } + + #[test] + fn parses_finished_event() { + let event = parse_event(msg( + r#"{"event_type":2,"result_code":3,"api_id":1100,"status":0}"#, + "", + )) + .unwrap(); + let ApiOperationEvent::Finished(result) = event else { + panic!("expected result"); + }; + assert_eq!(result.code, ApiOperationResultCode::Canceled); + assert!(!result.is_success()); + assert!(result.into_body().is_err()); + } + + #[test] + fn ignores_unknown_event_type() { + assert!(parse_event(msg(r#"{"event_type":0}"#, "")).is_none()); + assert!(parse_event(msg("not json", "")).is_none()); + } + + #[tokio::test] + async fn dispatch_delivers_result_and_closes_stream() { + let registry: Arc = Arc::new(Mutex::new(HashMap::new())); + let (sender, events) = mpsc::unbounded_channel(); + registry.lock().unwrap().insert("op-1".to_owned(), sender); + let mut operation = ApiOperation { + handle: ApiOperationHandle { + operation_id: "op-1".to_owned(), + api_id: 1100, + }, + events, + registry: Arc::clone(®istry), + }; + + dispatch_event(®istry, "t", msg(r#"{"event_type":1}"#, "p")); + dispatch_event( + ®istry, + "t", + msg(r#"{"event_type":2,"result_code":1,"status":0}"#, "done"), + ); + + let result = operation.wait().await.unwrap(); + assert_eq!(result.into_body().unwrap(), "done"); + assert!(operation.next_event().await.is_none()); + assert!(registry.lock().unwrap().is_empty()); + } +} diff --git a/booster_sdk/src/dds/rpc.rs b/booster_sdk/src/dds/rpc.rs index b7b9f49..4da3bf1 100644 --- a/booster_sdk/src/dds/rpc.rs +++ b/booster_sdk/src/dds/rpc.rs @@ -243,6 +243,30 @@ impl RpcClient { where R: DeserializeOwned + Send + 'static, { + let header = serde_json::json!({ "api_id": api_id }).to_string(); + let response_body = self + .call_with_header(api_id, header, body.into(), timeout) + .await?; + + decode_response_body(&response_body).map_err(|err| { + RpcError::RequestFailed { + status: 0, + message: format!("Failed to deserialize response body: {err}"), + } + .into() + }) + } + + /// Send a request with a pre-built JSON header and return the raw response body. + /// + /// Non-zero response statuses are returned as errors. + pub(crate) async fn call_with_header( + &self, + api_id: i32, + header: String, + body: String, + timeout: Option, + ) -> Result { if self.startup_wait > Duration::from_millis(0) && !self.startup_wait_done.swap(true, Ordering::SeqCst) { @@ -259,8 +283,6 @@ impl RpcClient { let mut response_stream = self.response_stream.lock().await; let request_id = Uuid::new_v4().to_string(); - let body = body.into(); - let header = serde_json::json!({ "api_id": api_id }).to_string(); let service_topic = self.service_topic.clone(); tracing::debug!( @@ -364,13 +386,7 @@ impl RpcClient { return Err(RpcError::from_status_code(status_code, message).into()); } - let result: R = - decode_response_body(&response.body).map_err(|err| RpcError::RequestFailed { - status: status_code, - message: format!("Failed to deserialize response body: {err}"), - })?; - - return Ok(result); + return Ok(response.body); } } } diff --git a/booster_sdk/src/dds/topics.rs b/booster_sdk/src/dds/topics.rs index ab8cd74..cc9d4e7 100644 --- a/booster_sdk/src/dds/topics.rs +++ b/booster_sdk/src/dds/topics.rs @@ -55,6 +55,9 @@ pub const X5_CAMERA_CONTROL_API_TOPIC: &str = "rt/X5CameraControl"; pub const CAMERA_API_TOPIC: &str = "rt/CameraApiTopic"; pub const HAND_EYE_CALIB_API_TOPIC: &str = "rt/HandEyeCalibApiTopic"; +pub const LOCO_API_OPERATION_EVENT_TOPIC: &str = "rt/LocoApiOperationEvent"; +pub const LUI_API_OPERATION_EVENT_TOPIC: &str = "rt/LuiApiOperationEvent"; + pub fn rpc_request_topic(service_topic: &str) -> TopicSpec { TopicSpec { name: format!("{service_topic}Req"), diff --git a/booster_sdk/src/types/error.rs b/booster_sdk/src/types/error.rs index e91bc93..6694479 100644 --- a/booster_sdk/src/types/error.rs +++ b/booster_sdk/src/types/error.rs @@ -81,6 +81,9 @@ pub enum RpcError { #[error("State transition failed: {0}")] StateTransitionFailed(String), + #[error("Motion request refused because the battery is low: {0}")] + LowBattery(String), + #[error("Invalid RPC status code: {0}")] InvalidStatusCode(i32), @@ -103,6 +106,7 @@ impl RpcError { 500 => RpcError::InternalServerError(message), 501 => RpcError::ServerRefused(message), 502 => RpcError::StateTransitionFailed(message), + 503 => RpcError::LowBattery(message), _ => RpcError::RequestFailed { status: code, message, diff --git a/booster_sdk_py/booster_sdk/client/lui.py b/booster_sdk_py/booster_sdk/client/lui.py index 2d0125e..60056c5 100644 --- a/booster_sdk_py/booster_sdk/client/lui.py +++ b/booster_sdk_py/booster_sdk/client/lui.py @@ -8,10 +8,18 @@ BoosterSdkError = bindings.BoosterSdkError LuiTtsConfig = bindings.LuiTtsConfig LuiTtsParameter = bindings.LuiTtsParameter +LuiSynthesizeSpeechRequest = bindings.LuiSynthesizeSpeechRequest +LuiSynthesizeSpeechResponse = bindings.LuiSynthesizeSpeechResponse +LuiRecognizeAudioRequest = bindings.LuiRecognizeAudioRequest +LuiRecognizeAudioResponse = bindings.LuiRecognizeAudioResponse __all__ = [ "LuiClient", "BoosterSdkError", "LuiTtsConfig", "LuiTtsParameter", + "LuiSynthesizeSpeechRequest", + "LuiSynthesizeSpeechResponse", + "LuiRecognizeAudioRequest", + "LuiRecognizeAudioResponse", ] diff --git a/booster_sdk_py/booster_sdk_bindings/booster_sdk_bindings.pyi b/booster_sdk_py/booster_sdk_bindings/booster_sdk_bindings.pyi index b6a6574..5eeaed0 100644 --- a/booster_sdk_py/booster_sdk_bindings/booster_sdk_bindings.pyi +++ b/booster_sdk_py/booster_sdk_bindings/booster_sdk_bindings.pyi @@ -687,10 +687,15 @@ class StartAiChatParameter: llm_config: LlmConfig, tts_config: TtsConfig, enable_face_tracking: bool, + persona_id: str | None = None, ) -> None: """Create AI chat startup parameters.""" ... @property + def persona_id(self) -> str | None: + """Optional AgentHub persona identifier.""" + ... + @property def interrupt_mode(self) -> bool: """Whether interruption mode is enabled.""" ... @@ -1118,6 +1123,10 @@ class BoosterClient: """Toggle upper-body custom control mode.""" ... + def reset_odometry_to(self, x: float, y: float, theta: float) -> None: + """Reset odometry to a target pose (meters, radians).""" + ... + def reset_odometry(self) -> None: """Reset base odometry estimate.""" ... @@ -1338,7 +1347,10 @@ class AudioCaptureStreamOptions: enable_naec_pcm: bool = ..., requested_raw_format: PcmFormat | None = ..., ) -> None: - """Create capture-stream initialization options.""" + """Create capture-stream initialization options. + + ``requested_raw_format`` defaults to 16 kHz, 3-channel, 16-bit PCM. + """ ... class InitPlayerResponse: @@ -1752,8 +1764,95 @@ class AiClient: """Disable AI face tracking mode.""" ... +class LuiSynthesizeSpeechRequest: + """Payload for :meth:`LuiClient.synthesize_speech`.""" + + def __init__( + self, + text: str, + voice_type: str = "default", + speed: float = 1.0, + playback: bool = False, + ) -> None: + """Create a synthesis request. + + Args: + text: Text to synthesize, at most 1000 Unicode code points. + voice_type: Voice identifier. + speed: Speech speed ratio; 1.0 is normal, 0.5 slowest, 2.0 fastest. + playback: Whether the robot should also play the audio. + """ + ... + @property + def text(self) -> str: ... + @property + def voice_type(self) -> str: ... + @property + def speed(self) -> float: ... + @property + def playback(self) -> bool: ... + +class LuiSynthesizeSpeechResponse: + """Audio returned by :meth:`LuiClient.synthesize_speech`.""" + + @property + def audio_base64(self) -> str: + """Base64-encoded audio data.""" + ... + @property + def sample_rate_hz(self) -> int: ... + @property + def channels(self) -> int: ... + @property + def bits_per_sample(self) -> int: ... + @property + def format(self) -> str: + """Audio format, e.g. ``pcm_s16le``.""" + ... + +class LuiRecognizeAudioRequest: + """Payload for the LUI audio recognition APIs.""" + + @staticmethod + def from_file(file_path: str) -> LuiRecognizeAudioRequest: + """Recognize a ``.wav`` or ``.mp3`` file on the robot.""" + ... + @staticmethod + def from_pcm_base64(audio_base64: str) -> LuiRecognizeAudioRequest: + """Recognize base64-encoded 16 kHz mono 16-bit little-endian raw PCM. + + Audio is limited to 120 s and 6 MiB decoded. + """ + ... + @property + def input_type(self) -> str: + """``"file"`` or ``"pcm"``.""" + ... + @property + def file_path(self) -> str: ... + @property + def sample_rate_hz(self) -> int: ... + @property + def channels(self) -> int: ... + @property + def bits_per_sample(self) -> int: ... + @property + def format(self) -> str: ... + +class LuiRecognizeAudioResponse: + """Recognition result.""" + + @property + def text(self) -> str: + """Recognized text.""" + ... + class LuiClient: - """Client for LUI ASR/TTS APIs.""" + """Client for LUI ASR/TTS APIs. + + ASR, TTS and audio recognizer sessions are owned by the client instance + that started them. + """ def __init__(self, startup_wait_sec: float | None = ...) -> None: """Create LUI client. @@ -1783,6 +1882,58 @@ class LuiClient: """Send text payload for TTS synthesis.""" ... + @property + def client_id(self) -> str: + """Identifier of this client instance.""" + ... + + @property + def current_asr_session_id(self) -> str | None: + """ASR session started by this client, if any.""" + ... + + @property + def current_audio_recognizer_session_id(self) -> str | None: + """Audio recognizer session started by this client, if any.""" + ... + + @property + def current_tts_session_id(self) -> str | None: + """TTS session started by this client, if any.""" + ... + + def synthesize_speech( + self, req: LuiSynthesizeSpeechRequest + ) -> LuiSynthesizeSpeechResponse: + """Synthesize speech in this client's TTS session. + + Requires :meth:`start_tts`. + """ + ... + + def recognize_audio_once( + self, req: LuiRecognizeAudioRequest + ) -> LuiRecognizeAudioResponse: + """Recognize an audio file or PCM payload without a session.""" + ... + + def start_audio_recognizer(self) -> None: + """Start a reusable audio recognizer session owned by this client.""" + ... + + def stop_audio_recognizer(self) -> None: + """Stop this client's audio recognizer session.""" + ... + + def recognize_audio_in_session( + self, req: LuiRecognizeAudioRequest + ) -> LuiRecognizeAudioResponse: + """Recognize audio in this client's recognizer session. + + Requires :meth:`start_audio_recognizer`. + """ + ... + class LightControlClient: """Client for LED light control APIs.""" diff --git a/booster_sdk_py/src/client/ai.rs b/booster_sdk_py/src/client/ai.rs index 2ef9028..87ef713 100644 --- a/booster_sdk_py/src/client/ai.rs +++ b/booster_sdk_py/src/client/ai.rs @@ -113,14 +113,17 @@ pub struct PyStartAiChatParameter(StartAiChatParameter); #[pymethods] impl PyStartAiChatParameter { #[new] + #[pyo3(signature = (interrupt_mode, asr_config, llm_config, tts_config, enable_face_tracking, persona_id=None))] fn new( interrupt_mode: bool, asr_config: PyAsrConfig, llm_config: PyLlmConfig, tts_config: PyTtsConfig, enable_face_tracking: bool, + persona_id: Option, ) -> Self { Self(StartAiChatParameter { + persona_id, interrupt_mode, asr_config: asr_config.into(), llm_config: llm_config.into(), @@ -129,6 +132,11 @@ impl PyStartAiChatParameter { }) } + #[getter] + fn persona_id(&self) -> Option { + self.0.persona_id.clone() + } + #[getter] fn interrupt_mode(&self) -> bool { self.0.interrupt_mode diff --git a/booster_sdk_py/src/client/audio.rs b/booster_sdk_py/src/client/audio.rs index bfb6bab..e637dc4 100644 --- a/booster_sdk_py/src/client/audio.rs +++ b/booster_sdk_py/src/client/audio.rs @@ -381,7 +381,10 @@ impl PyAudioCaptureStreamOptions { Self(AudioCaptureStreamOptions { enable_raw_pcm, enable_naec_pcm, - requested_raw_format: requested_raw_format.map(Into::into).unwrap_or_default(), + requested_raw_format: requested_raw_format.map_or_else( + || AudioCaptureStreamOptions::default().requested_raw_format, + Into::into, + ), }) } } diff --git a/booster_sdk_py/src/client/booster.rs b/booster_sdk_py/src/client/booster.rs index a85719e..43706c0 100644 --- a/booster_sdk_py/src/client/booster.rs +++ b/booster_sdk_py/src/client/booster.rs @@ -1818,6 +1818,15 @@ impl PyBoosterClient { wait_for_future(py, async move { client.reset_odometry().await }).map_err(to_py_err) } + fn reset_odometry_to(&self, py: Python<'_>, x: f64, y: f64, theta: f64) -> PyResult<()> { + let client = Arc::clone(&self.client); + wait_for_future( + py, + async move { client.reset_odometry_to(x, y, theta).await }, + ) + .map_err(to_py_err) + } + fn load_custom_trained_traj( &self, py: Python<'_>, diff --git a/booster_sdk_py/src/client/lui.rs b/booster_sdk_py/src/client/lui.rs index 85ad4dc..8021f70 100644 --- a/booster_sdk_py/src/client/lui.rs +++ b/booster_sdk_py/src/client/lui.rs @@ -1,6 +1,9 @@ use std::sync::Arc; -use booster_sdk::client::ai::{LuiClient, LuiTtsConfig, LuiTtsParameter}; +use booster_sdk::client::ai::{ + LuiClient, LuiRecognizeAudioRequest, LuiRecognizeAudioResponse, LuiSynthesizeSpeechRequest, + LuiSynthesizeSpeechResponse, LuiTtsConfig, LuiTtsParameter, +}; use pyo3::{Bound, prelude::*, types::PyModule}; use crate::{runtime::wait_for_future, startup_wait_from_seconds, to_py_err}; @@ -51,6 +54,135 @@ impl From for LuiTtsParameter { } } +#[pyclass(module = "booster_sdk_bindings", name = "LuiSynthesizeSpeechRequest")] +#[derive(Clone)] +pub struct PyLuiSynthesizeSpeechRequest(LuiSynthesizeSpeechRequest); + +#[pymethods] +impl PyLuiSynthesizeSpeechRequest { + #[new] + #[pyo3(signature = (text, voice_type="default".to_owned(), speed=1.0, playback=false))] + fn new(text: String, voice_type: String, speed: f64, playback: bool) -> Self { + Self(LuiSynthesizeSpeechRequest { + text, + voice_type, + speed, + playback, + }) + } + + #[getter] + fn text(&self) -> String { + self.0.text.clone() + } + + #[getter] + fn voice_type(&self) -> String { + self.0.voice_type.clone() + } + + #[getter] + fn speed(&self) -> f64 { + self.0.speed + } + + #[getter] + fn playback(&self) -> bool { + self.0.playback + } +} + +#[pyclass(module = "booster_sdk_bindings", name = "LuiSynthesizeSpeechResponse")] +#[derive(Clone)] +pub struct PyLuiSynthesizeSpeechResponse(LuiSynthesizeSpeechResponse); + +#[pymethods] +impl PyLuiSynthesizeSpeechResponse { + #[getter] + fn audio_base64(&self) -> String { + self.0.audio_base64.clone() + } + + #[getter] + fn sample_rate_hz(&self) -> i32 { + self.0.sample_rate_hz + } + + #[getter] + fn channels(&self) -> i32 { + self.0.channels + } + + #[getter] + fn bits_per_sample(&self) -> i32 { + self.0.bits_per_sample + } + + #[getter] + fn format(&self) -> String { + self.0.format.clone() + } +} + +#[pyclass(module = "booster_sdk_bindings", name = "LuiRecognizeAudioRequest")] +#[derive(Clone)] +pub struct PyLuiRecognizeAudioRequest(LuiRecognizeAudioRequest); + +#[pymethods] +impl PyLuiRecognizeAudioRequest { + #[staticmethod] + fn from_file(file_path: String) -> Self { + Self(LuiRecognizeAudioRequest::from_file(file_path)) + } + + #[staticmethod] + fn from_pcm_base64(audio_base64: String) -> Self { + Self(LuiRecognizeAudioRequest::from_pcm_base64(audio_base64)) + } + + #[getter] + fn input_type(&self) -> String { + self.0.input_type.clone() + } + + #[getter] + fn file_path(&self) -> String { + self.0.file_path.clone() + } + + #[getter] + fn sample_rate_hz(&self) -> i32 { + self.0.sample_rate_hz + } + + #[getter] + fn channels(&self) -> i32 { + self.0.channels + } + + #[getter] + fn bits_per_sample(&self) -> i32 { + self.0.bits_per_sample + } + + #[getter] + fn format(&self) -> String { + self.0.format.clone() + } +} + +#[pyclass(module = "booster_sdk_bindings", name = "LuiRecognizeAudioResponse")] +#[derive(Clone)] +pub struct PyLuiRecognizeAudioResponse(LuiRecognizeAudioResponse); + +#[pymethods] +impl PyLuiRecognizeAudioResponse { + #[getter] + fn text(&self) -> String { + self.0.text.clone() + } +} + #[pyclass(module = "booster_sdk_bindings", name = "LuiClient", unsendable)] pub struct PyLuiClient { client: Arc, @@ -99,11 +231,80 @@ impl PyLuiClient { let param = param.into(); wait_for_future(py, async move { client.send_tts_text(¶m).await }).map_err(to_py_err) } + + #[getter] + fn client_id(&self) -> String { + self.client.client_id().to_owned() + } + + #[getter] + fn current_asr_session_id(&self) -> Option { + self.client.current_asr_session_id() + } + + #[getter] + fn current_audio_recognizer_session_id(&self) -> Option { + self.client.current_audio_recognizer_session_id() + } + + #[getter] + fn current_tts_session_id(&self) -> Option { + self.client.current_tts_session_id() + } + + fn synthesize_speech( + &self, + py: Python<'_>, + req: PyLuiSynthesizeSpeechRequest, + ) -> PyResult { + let client = Arc::clone(&self.client); + wait_for_future(py, async move { client.synthesize_speech(&req.0).await }) + .map(PyLuiSynthesizeSpeechResponse) + .map_err(to_py_err) + } + + fn recognize_audio_once( + &self, + py: Python<'_>, + req: PyLuiRecognizeAudioRequest, + ) -> PyResult { + let client = Arc::clone(&self.client); + wait_for_future(py, async move { client.recognize_audio_once(&req.0).await }) + .map(PyLuiRecognizeAudioResponse) + .map_err(to_py_err) + } + + fn start_audio_recognizer(&self, py: Python<'_>) -> PyResult<()> { + let client = Arc::clone(&self.client); + wait_for_future(py, async move { client.start_audio_recognizer().await }).map_err(to_py_err) + } + + fn stop_audio_recognizer(&self, py: Python<'_>) -> PyResult<()> { + let client = Arc::clone(&self.client); + wait_for_future(py, async move { client.stop_audio_recognizer().await }).map_err(to_py_err) + } + + fn recognize_audio_in_session( + &self, + py: Python<'_>, + req: PyLuiRecognizeAudioRequest, + ) -> PyResult { + let client = Arc::clone(&self.client); + wait_for_future(py, async move { + client.recognize_audio_in_session(&req.0).await + }) + .map(PyLuiRecognizeAudioResponse) + .map_err(to_py_err) + } } pub(crate) fn register(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; m.add_class::()?; Ok(()) } From eaddf5cb2c70ba5a67346b6e5a5bc38293465012 Mon Sep 17 00:00:00 2001 From: Gijs de Jong Date: Wed, 23 Sep 2026 13:56:32 +0200 Subject: [PATCH 2/3] version 0.1.3 --- Cargo.lock | 12 ++++++------ Cargo.toml | 2 +- examples/rust/locomotion/Cargo.toml | 2 +- examples/rust/look_around/Cargo.toml | 2 +- pixi.toml | 2 +- 5 files changed, 10 insertions(+), 10 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 55d4884..ae386fa 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -52,7 +52,7 @@ checksum = "c4512299f36f043ab09a583e57bceb5a5aab7a73db1805848e8fef3c9e8c78b3" [[package]] name = "booster_sdk" -version = "0.1.2" +version = "0.1.3" dependencies = [ "futures", "rustdds", @@ -68,7 +68,7 @@ dependencies = [ [[package]] name = "booster_sdk_py" -version = "0.1.2" +version = "0.1.3" dependencies = [ "booster_sdk", "pyo3", @@ -491,7 +491,7 @@ dependencies = [ [[package]] name = "iana-time-zone-haiku" -version = "0.1.2" +version = "0.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f31827a206f56af32e590ba56d5d2d085f558508192593743f16b2306495269f" dependencies = [ @@ -653,7 +653,7 @@ dependencies = [ [[package]] name = "locomotion" -version = "0.1.2" +version = "0.1.3" dependencies = [ "booster_sdk", "tokio", @@ -669,7 +669,7 @@ checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" [[package]] name = "look_around" -version = "0.1.2" +version = "0.1.3" dependencies = [ "booster_sdk", "tokio", @@ -1312,7 +1312,7 @@ dependencies = [ [[package]] name = "serde_repr" -version = "0.1.20" +version = "0.1.30" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "175ee3e80ae9982737ca543e96133087cbd9a485eecc3bc4de9c1a37b47ea59c" dependencies = [ diff --git a/Cargo.toml b/Cargo.toml index 885e7ca..dcb8bf8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -3,7 +3,7 @@ members = ["booster_sdk", "booster_sdk_py", "examples/rust/*"] resolver = "2" [workspace.package] -version = "0.1.2" +version = "0.1.3" edition = "2024" authors = ["Team whIRLwind"] license = "MIT OR Apache-2.0" diff --git a/examples/rust/locomotion/Cargo.toml b/examples/rust/locomotion/Cargo.toml index 50af39d..2454f65 100644 --- a/examples/rust/locomotion/Cargo.toml +++ b/examples/rust/locomotion/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "locomotion" -version = "0.1.2" +version = "0.1.3" edition = "2024" [dependencies] diff --git a/examples/rust/look_around/Cargo.toml b/examples/rust/look_around/Cargo.toml index d6a605f..f61af75 100644 --- a/examples/rust/look_around/Cargo.toml +++ b/examples/rust/look_around/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "look_around" -version = "0.1.2" +version = "0.1.3" edition = "2024" [dependencies] diff --git a/pixi.toml b/pixi.toml index b182898..6f37b56 100644 --- a/pixi.toml +++ b/pixi.toml @@ -3,7 +3,7 @@ authors = ["Team whIRLwind"] channels = ["conda-forge"] name = "booster-sdk" platforms = ["osx-arm64", "linux-64", "linux-aarch64"] -version = "0.1.2" +version = "0.1.3" [environments] py = ["wheel-build", "python-tasks"] From 0218f11b953a7568feebceb2cdd811c5f3d202a1 Mon Sep 17 00:00:00 2001 From: Gijs de Jong Date: Wed, 23 Sep 2026 14:02:39 +0200 Subject: [PATCH 3/3] update versions --- Cargo.lock | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index ae386fa..431a81e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -491,7 +491,7 @@ dependencies = [ [[package]] name = "iana-time-zone-haiku" -version = "0.1.3" +version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f31827a206f56af32e590ba56d5d2d085f558508192593743f16b2306495269f" dependencies = [ @@ -1312,7 +1312,7 @@ dependencies = [ [[package]] name = "serde_repr" -version = "0.1.30" +version = "0.1.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "175ee3e80ae9982737ca543e96133087cbd9a485eecc3bc4de9c1a37b47ea59c" dependencies = [