返回 CodeWhale
speech_host_tests.rs
根目录 / crates / tui / src / tools / speech_host_tests.rs
1 //! Actual CLI/ToolSpec consumers, pinned host and loopback provider. No live provider calls.
2 use super::*;
3 use crate::config::Config;
4 use crate::extension_host::tests::node_for_tests;
5 use crate::extension_host::{ExtensionHostManager, ExtensionHostOptions, TestManagerGuard};
6 use crate::features::{Feature, Features, FeaturesToml};
7 use crate::plugins::activation::TestPolicyGuard;
8 use std::sync::Arc;
9 use wiremock::matchers::{method, path};
10 use wiremock::{Mock, MockServer, ResponseTemplate};
11
12 pub(crate) fn manager(node: PathBuf, home: &Path) -> Arc<ExtensionHostManager> {
13 Arc::new(ExtensionHostManager::new(ExtensionHostOptions {
14 runtime: crate::config::ExtensionHostRuntime::Node,
15 node_override: Some(node),
16 root: Some(home.join("host")),
17 ..ExtensionHostOptions::default()
18 }))
19 }
20 fn context(root: &Path, host: bool) -> ToolContext {
21 let mut features = Features::with_defaults();
22 if host {
23 features.enable(Feature::SpeechHost);
24 }
25 ToolContext::new(root).with_features(features)
26 }
27 pub(crate) fn config(url: &str, host: bool) -> Config {
28 let mut config = Config {
29 provider: Some("xiaomi-mimo".into()),
30 ..Config::default()
31 };
32 config
33 .set_provider_base_url_override(
34 &config.test_identity_for_kind(ProviderKind::XiaomiMimo),
35 Some(url.into()),
36 )
37 .unwrap();
38 config
39 .set_provider_api_key_override(
40 &config.test_identity_for_kind(ProviderKind::XiaomiMimo),
41 Some("local-fixture-only".into()),
42 )
43 .unwrap();
44 config.features = Some(FeaturesToml {
45 entries: [("speech_host".into(), host)].into_iter().collect(),
46 });
47 config
48 }
49 pub(crate) async fn provider(server: &MockServer, count: u64) {
50 Mock::given(method("POST"))
51 .and(path("/v1/chat/completions"))
52 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
53 "choices":[{"message":{"audio":{"data":"aGk=","transcript":"hi"}}}]
54 })))
55 .expect(count)
56 .mount(server)
57 .await;
58 }
59 fn equal_results(rust: &ToolResult, host: &ToolResult) {
60 assert_eq!(host.success, rust.success);
61 assert_eq!(
62 host.content, rust.content,
63 "exact model-visible result bytes"
64 );
65 assert_eq!(host.metadata, rust.metadata);
66 }
67
68 #[test]
69 fn speech_host_is_independent_and_defaults_to_rust() {
70 let flags = Features::with_defaults();
71 assert!(!flags.enabled(Feature::SpeechHost));
72 assert_eq!(
73 crate::features::feature_from_key("speech_host"),
74 Some(Feature::SpeechHost)
75 );
76 let mut flags = flags;
77 flags.enable(Feature::SpeechHost);
78 assert!(!flags.enabled(Feature::FinanceHost));
79 assert!(!flags.enabled(Feature::DataHost));
80 assert!(!flags.enabled(Feature::ExtensionHost));
81 }
82
83 #[tokio::test(flavor = "current_thread")]
84 async fn real_host_speech_options_and_formats_match_both_legacy_surfaces() {
85 let _home = crate::test_support::SealedHome::new();
86 let _policy = TestPolicyGuard::extension_host(false);
87 let Some(node) = node_for_tests("real_host_speech_options") else {
88 return;
89 };
90 let home = tempfile::tempdir().unwrap();
91 let manager = manager(node, home.path());
92 let _manager = TestManagerGuard::install(Arc::clone(&manager));
93 for surface in [SpeechSurface::Tool, SpeechSurface::Cli] {
94 let cases = [
95 (None, None, None, None, false),
96 (Some("mimo-tts"), Some("Mia"), Some(" warm "), None, false),
97 (
98 None,
99 None,
100 Some(" slow "),
101 Some("\u{85}Bright\u{3000}"),
102 false,
103 ),
104 (
105 None,
106 None,
107 Some("\u{feff}calm\u{feff}"),
108 Some("\u{2007}"),
109 false,
110 ),
111 (
112 None,
113 Some("data:audio/wav;base64,c2FtcGxl"),
114 None,
115 None,
116 false,
117 ),
118 (None, None, None, None, true),
119 (Some("mimo-chat"), None, None, None, false),
120 (Some("mimo-v2.5-tts-voiceclone"), None, None, None, false),
121 (Some("mimo-v2.5-tts-voicedesign"), None, None, None, false),
122 (None, Some(""), None, None, true),
123 (None, Some(""), None, None, false),
124 (
125 Some("mimo-v2.5-tts-voicedesign"),
126 None,
127 Some("warm"),
128 None,
129 true,
130 ),
131 ];
132 for (model, voice, instruction, voice_prompt, has_clone_path) in cases {
133 let inputs = || SpeechPreparation {
134 model,
135 voice,
136 instruction: instruction.map(str::to_string),
137 voice_prompt: voice_prompt.map(str::to_string),
138 has_clone_path,
139 surface,
140 };
141 let rust = prepare_speech_options(inputs(), &context(home.path(), false)).await;
142 let host = prepare_speech_options(inputs(), &context(home.path(), true)).await;
143 match (rust, host) {
144 (Ok(rust), Ok(host)) => assert_eq!(host, rust),
145 (Err(rust), Err(host)) => assert_eq!(host.to_string(), rust.to_string()),
146 values => panic!("speech preparation parity mismatch: {values:?}"),
147 }
148 }
149 for format in [
150 "WAV",
151 " pcm ",
152 "\u{85}MP3\u{85}",
153 "pcm16",
154 "flac",
155 "\u{feff}wav",
156 ] {
157 let rust = prepare_speech_format(format, surface, &context(home.path(), false)).await;
158 let host = prepare_speech_format(format, surface, &context(home.path(), true)).await;
159 assert_eq!(
160 host.map_err(|error| error.to_string()),
161 rust.map_err(|error| error.to_string())
162 );
163 }
164 }
165 assert!(!crate::plugins::activation::extension_host_policy_enabled());
166 manager.shutdown().await;
167 }
168
169 #[tokio::test(flavor = "current_thread")]
170 async fn real_host_speech_tool_and_hidden_alias_match_requests_results_and_audio() {
171 let _home = crate::test_support::SealedHome::new();
172 let _policy = TestPolicyGuard::extension_host(false);
173 let Some(node) = node_for_tests("real_host_speech_tool") else {
174 return;
175 };
176 let home = tempfile::tempdir().unwrap();
177 let manager = manager(node, home.path());
178 let _manager = TestManagerGuard::install(Arc::clone(&manager));
179 let server = MockServer::start().await;
180 let client = CodewhaleClient::new(&config(&server.uri(), false)).unwrap();
181 std::fs::write(home.path().join("sample.wav"), b"sample").unwrap();
182 let inputs = [
183 json!({"text":" hello ","model":"mimo-tts","format":"pcm","output":"audio/result.pcm16"}),
184 json!({"text":"hello","voice_prompt":"\u{85}Bright\u{3000}","instruction":" slow ","output":"audio/result.wav"}),
185 json!({"text":"hello","voice":"data:audio/wav;base64,c2FtcGxl","instruction":"clone","output":"audio/result.wav"}),
186 json!({"text":"hello","clone_voice":"sample.wav","output":"audio/result.wav"}),
187 json!({"text":"hello","model":"mimo-v2.5-tts-voicedesign","clone_voice":"sample.wav","instruction":"warm","output":"audio/result.wav"}),
188 ];
189 for name in ["speech", "tts"] {
190 let tool = if name == "tts" {
191 SpeechTool::alias(name, Some(client.clone()), None)
192 } else {
193 SpeechTool::new(name, Some(client.clone()), None)
194 };
195 assert_eq!(tool.model_visible(), name == "speech");
196 for input in &inputs {
197 server.reset().await;
198 provider(&server, 2).await;
199 let expected = tool
200 .execute(input.clone(), &context(home.path(), false))
201 .await
202 .unwrap();
203 let actual = tool
204 .execute(input.clone(), &context(home.path(), true))
205 .await
206 .unwrap();
207 equal_results(&expected, &actual);
208 assert_eq!(
209 std::fs::read(home.path().join(input["output"].as_str().unwrap())).unwrap(),
210 b"hi"
211 );
212 let requests = server.received_requests().await.unwrap();
213 assert_eq!(requests.len(), 2);
214 assert_eq!(
215 requests[0].body, requests[1].body,
216 "exact provider request bytes"
217 );
218 }
219 }
220 manager.shutdown().await;
221 }
222
223 #[tokio::test(flavor = "current_thread")]
224 async fn speech_host_refusal_and_cancel_never_fallback_or_write() {
225 let _home = crate::test_support::SealedHome::new();
226 let _policy = TestPolicyGuard::extension_host(false);
227 let home = tempfile::tempdir().unwrap();
228 let server = MockServer::start().await;
229 let client = CodewhaleClient::new(&config(&server.uri(), false)).unwrap();
230 let manager = manager(home.path().join("missing-node"), home.path());
231 let _manager = TestManagerGuard::install(Arc::clone(&manager));
232 let tool = SpeechTool::new("speech", Some(client), None);
233 let input =
234 json!({"text":"hello","clone_voice":"unread-missing.wav","output":"must-not-exist.wav"});
235 let error = tool
236 .execute(input.clone(), &context(home.path(), true))
237 .await
238 .unwrap_err();
239 assert!(!matches!(error, ToolError::InvalidInput { .. }));
240 assert!(server.received_requests().await.unwrap().is_empty());
241 assert!(!home.path().join("must-not-exist.wav").exists());
242 let token = tokio_util::sync::CancellationToken::new();
243 token.cancel();
244 let error = tool
245 .execute(input, &context(home.path(), true).with_cancel_token(token))
246 .await
247 .unwrap_err();
248 assert!(error.to_string().contains("cancel"));
249 manager.shutdown().await;
250 }
251
252 #[tokio::test(flavor = "current_thread")]
253 async fn speech_tool_preserves_validation_and_file_error_order() {
254 let _home = crate::test_support::SealedHome::new();
255 let _policy = TestPolicyGuard::extension_host(false);
256 let Some(node) = node_for_tests("speech_error_order") else {
257 return;
258 };
259 let home = tempfile::tempdir().unwrap();
260 let manager = manager(node, home.path());
261 let _manager = TestManagerGuard::install(Arc::clone(&manager));
262 let server = MockServer::start().await;
263 let tool = SpeechTool::new(
264 "speech",
265 Some(CodewhaleClient::new(&config(&server.uri(), false)).unwrap()),
266 None,
267 );
268 for input in [
269 json!({"text":"hello","format":"flac","output":"../escape.wav","model":"mimo-chat"}),
270 json!({"text":"hello","output":"../escape.wav","model":"mimo-chat"}),
271 json!({"text":"hello","clone_voice":"missing.wav","format":"pcm"}),
272 json!({"text":"hello","voice_prompt":"","model":"mimo-v2.5-tts-voicedesign"}),
273 ] {
274 let rust = tool
275 .execute(input.clone(), &context(home.path(), false))
276 .await
277 .unwrap_err();
278 let host = tool
279 .execute(input, &context(home.path(), true))
280 .await
281 .unwrap_err();
282 assert_eq!(host.to_string(), rust.to_string());
283 }
284 manager.shutdown().await;
285 }
286
287 #[tokio::test(flavor = "current_thread")]
288 async fn real_host_speech_retains_core_network_policy_before_provider_calls() {
289 use crate::network_policy::{NetworkPolicy, NetworkPolicyDecider};
290 let _home = crate::test_support::SealedHome::new();
291 let _policy = TestPolicyGuard::extension_host(false);
292 let Some(node) = node_for_tests("speech_network_policy") else {
293 return;
294 };
295 let home = tempfile::tempdir().unwrap();
296 let manager = manager(node, home.path());
297 let _manager = TestManagerGuard::install(Arc::clone(&manager));
298 let server = MockServer::start().await;
299 let tool = SpeechTool::new(
300 "speech",
301 Some(CodewhaleClient::new(&config(&server.uri(), false)).unwrap()),
302 None,
303 );
304 for decision in [Decision::Deny, Decision::Prompt] {
305 let policy = || {
306 NetworkPolicyDecider::new(
307 NetworkPolicy {
308 default: decision.into(),
309 allow: Vec::new(),
310 deny: Vec::new(),
311 proxy: Vec::new(),
312 proxy_fake_ip_cidrs: Vec::new(),
313 audit: false,
314 },
315 None,
316 )
317 };
318 let input = json!({"text":"hello","output":"blocked.wav"});
319 let rust = tool
320 .execute(
321 input.clone(),
322 &context(home.path(), false).with_network_policy(policy()),
323 )
324 .await
325 .unwrap_err();
326 let host = tool
327 .execute(
328 input,
329 &context(home.path(), true).with_network_policy(policy()),
330 )
331 .await
332 .unwrap_err();
333 assert_eq!(host.to_string(), rust.to_string());
334 }
335 assert!(server.received_requests().await.unwrap().is_empty());
336 assert!(!home.path().join("blocked.wav").exists());
337 manager.shutdown().await;
338 }
339
339 lines RUST