1use std::collections::HashMap;
2use std::fmt;
3use std::path::{Path, PathBuf};
4use std::time::Instant;
5
6use reqwest::header::{HeaderMap, CONTENT_TYPE};
7use serde_json::{json, Value};
8
9use crate::{
10 AgentProvider, CancellationToken, InvocationData, ProviderEventSink, ProviderFuture,
11 ProviderRequest, ProviderResponse, ProviderStage, RuntimeError,
12};
13
14const PROVIDER_PROBE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
15
16pub struct BinaryProviderResponse {
17 pub content_type: String,
18 pub body: Vec<u8>,
19}
20
21#[derive(Clone, PartialEq)]
23pub enum HttpCapabilityRoute {
24 OpenAiChat {
25 model: String,
26 persona: Value,
27 },
28 OpenAiEmbedding {
29 model: String,
30 },
31 ElevenLabsSpeech {
32 voice_id: String,
33 },
34 OpenAiTranscription {
35 model: String,
36 file_name: String,
37 content_type: String,
38 },
39 #[cfg(feature = "local-whisper")]
40 LocalWhisper {
41 model_path: PathBuf,
42 language: Option<String>,
43 },
44}
45
46impl fmt::Debug for HttpCapabilityRoute {
47 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
48 match self {
49 Self::OpenAiChat { model, .. } => formatter
50 .debug_struct("OpenAiChat")
51 .field("model", model)
52 .field("persona", &"[REDACTED]")
53 .finish(),
54 Self::OpenAiEmbedding { model } => formatter
55 .debug_struct("OpenAiEmbedding")
56 .field("model", model)
57 .finish(),
58 Self::ElevenLabsSpeech { voice_id } => formatter
59 .debug_struct("ElevenLabsSpeech")
60 .field("voice_id", voice_id)
61 .finish(),
62 Self::OpenAiTranscription {
63 model,
64 file_name,
65 content_type,
66 } => formatter
67 .debug_struct("OpenAiTranscription")
68 .field("model", model)
69 .field("file_name", file_name)
70 .field("content_type", content_type)
71 .finish(),
72 #[cfg(feature = "local-whisper")]
73 Self::LocalWhisper { language, .. } => formatter
74 .debug_struct("LocalWhisper")
75 .field("model_path", &"[REDACTED]")
76 .field("language", language)
77 .finish(),
78 }
79 }
80}
81
82pub struct HttpCapabilityProvider {
87 name: String,
88 base_url: String,
89 token: Option<String>,
90 routes: HashMap<String, HttpCapabilityRoute>,
91}
92
93impl fmt::Debug for HttpCapabilityProvider {
94 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
95 formatter
96 .debug_struct("HttpCapabilityProvider")
97 .field("name", &self.name)
98 .field("base_url", &self.base_url)
99 .field("token", &self.token.as_ref().map(|_| "[REDACTED]"))
100 .field("capabilities", &self.routes.keys().collect::<Vec<_>>())
101 .finish()
102 }
103}
104
105impl HttpCapabilityProvider {
106 pub fn new(
107 name: impl Into<String>,
108 base_url: impl Into<String>,
109 token: Option<String>,
110 ) -> Result<Self, RuntimeError> {
111 let name = name.into();
112 if name.trim().is_empty() {
113 return Err(RuntimeError::InvalidDefinition(
114 "provider name is required".to_string(),
115 ));
116 }
117 let base_url = base_url.into();
118 provider_url(&base_url, "models").map_err(RuntimeError::InvalidDefinition)?;
119 Ok(Self {
120 name,
121 base_url,
122 token: token
123 .map(|value| value.trim().to_string())
124 .filter(|value| !value.is_empty()),
125 routes: HashMap::new(),
126 })
127 }
128
129 pub fn local(name: impl Into<String>) -> Result<Self, RuntimeError> {
130 let name = name.into();
131 if name.trim().is_empty() {
132 return Err(RuntimeError::InvalidDefinition(
133 "provider name is required".to_string(),
134 ));
135 }
136 Ok(Self {
137 name,
138 base_url: String::new(),
139 token: None,
140 routes: HashMap::new(),
141 })
142 }
143
144 pub fn add_route(
145 &mut self,
146 capability: impl Into<String>,
147 route: HttpCapabilityRoute,
148 ) -> Result<(), RuntimeError> {
149 let capability = capability.into().trim().to_ascii_lowercase();
150 if capability.is_empty() || capability.len() > 128 {
151 return Err(RuntimeError::InvalidDefinition(
152 "provider capability is invalid".to_string(),
153 ));
154 }
155 self.routes.insert(capability, route);
156 Ok(())
157 }
158
159 pub fn with_route(
160 mut self,
161 capability: impl Into<String>,
162 route: HttpCapabilityRoute,
163 ) -> Result<Self, RuntimeError> {
164 self.add_route(capability, route)?;
165 Ok(self)
166 }
167}
168
169impl AgentProvider for HttpCapabilityProvider {
170 fn supports(&self, capability: &str) -> bool {
171 self.routes.contains_key(capability)
172 }
173
174 fn invoke<'a>(
175 &'a self,
176 request: ProviderRequest,
177 cancellation: CancellationToken,
178 ) -> ProviderFuture<'a> {
179 self.invoke_with_events(request, cancellation, ProviderEventSink::discard())
180 }
181
182 fn invoke_with_events<'a>(
183 &'a self,
184 request: ProviderRequest,
185 cancellation: CancellationToken,
186 events: ProviderEventSink,
187 ) -> ProviderFuture<'a> {
188 Box::pin(async move {
189 let route = self.routes.get(&request.capability).ok_or_else(|| {
190 RuntimeError::CapabilityUnavailable {
191 provider: self.name.clone(),
192 capability: request.capability.clone(),
193 }
194 })?;
195 if cancellation.is_cancelled() {
196 return Err(RuntimeError::Cancelled);
197 }
198 let response = match route {
199 HttpCapabilityRoute::OpenAiChat { model, persona } => {
200 let InvocationData::Json(payload) = &request.data else {
201 return Err(RuntimeError::InvalidDefinition(
202 "chat capability requires JSON input".to_string(),
203 ));
204 };
205 let data = provider_json_result(
206 &self.name,
207 &events,
208 "chat",
209 openai_chat_completion_result(
210 &self.base_url,
211 self.token.as_deref(),
212 model,
213 payload,
214 persona,
215 )
216 .await,
217 )?;
218 if cancellation.is_cancelled() {
219 return Err(RuntimeError::Cancelled);
220 }
221 validate_with_events(&self.name, &events, "chat", || {
222 validate_openai_chat_response(&data)
223 })?;
224 ProviderResponse {
225 data: InvocationData::Json(data),
226 metadata: json!({ "contentType": "application/json" }),
227 state: None,
228 }
229 }
230 HttpCapabilityRoute::OpenAiEmbedding { model } => {
231 let InvocationData::Json(payload) = &request.data else {
232 return Err(RuntimeError::InvalidDefinition(
233 "embedding capability requires JSON input".to_string(),
234 ));
235 };
236 let data = provider_json_result(
237 &self.name,
238 &events,
239 "embedding",
240 openai_embeddings_result(
241 &self.base_url,
242 self.token.as_deref(),
243 model,
244 payload,
245 )
246 .await,
247 )?;
248 if cancellation.is_cancelled() {
249 return Err(RuntimeError::Cancelled);
250 }
251 validate_with_events(&self.name, &events, "embedding", || {
252 validate_openai_embedding_response(&data)
253 })?;
254 ProviderResponse {
255 data: InvocationData::Json(data),
256 metadata: json!({ "contentType": "application/json" }),
257 state: None,
258 }
259 }
260 HttpCapabilityRoute::ElevenLabsSpeech { voice_id } => {
261 let InvocationData::Json(payload) = &request.data else {
262 return Err(RuntimeError::InvalidDefinition(
263 "speech capability requires JSON input".to_string(),
264 ));
265 };
266 let response =
267 elevenlabs_speech(&self.base_url, self.token.as_deref(), voice_id, payload)
268 .await
269 .map_err(|message| RuntimeError::provider(&self.name, message))?;
270 ProviderResponse {
271 data: InvocationData::Binary(response.body),
272 metadata: json!({ "contentType": response.content_type }),
273 state: None,
274 }
275 }
276 HttpCapabilityRoute::OpenAiTranscription {
277 model,
278 file_name,
279 content_type,
280 } => {
281 let InvocationData::Binary(audio) = &request.data else {
282 return Err(RuntimeError::InvalidDefinition(
283 "transcription capability requires binary input".to_string(),
284 ));
285 };
286 let data = provider_json_result(
287 &self.name,
288 &events,
289 "transcription",
290 openai_audio_transcription_result(
291 &self.base_url,
292 self.token.as_deref(),
293 model,
294 audio.clone(),
295 file_name,
296 content_type,
297 )
298 .await,
299 )?;
300 if cancellation.is_cancelled() {
301 return Err(RuntimeError::Cancelled);
302 }
303 validate_with_events(&self.name, &events, "transcription", || {
304 validate_transcription_response(&data)
305 })?;
306 ProviderResponse {
307 data: InvocationData::Json(data),
308 metadata: json!({ "contentType": "application/json" }),
309 state: None,
310 }
311 }
312 #[cfg(feature = "local-whisper")]
313 HttpCapabilityRoute::LocalWhisper {
314 model_path,
315 language,
316 } => {
317 let InvocationData::Binary(audio) = &request.data else {
318 return Err(RuntimeError::InvalidDefinition(
319 "transcription capability requires binary input".to_string(),
320 ));
321 };
322 let data = json!({
323 "text": local_whisper_transcription(
324 model_path,
325 audio,
326 request
327 .metadata
328 .pointer("/binding/language")
329 .and_then(Value::as_str)
330 .or(language.as_deref()),
331 )
332 .map_err(|message| RuntimeError::provider(&self.name, message))?,
333 });
334 if cancellation.is_cancelled() {
335 return Err(RuntimeError::Cancelled);
336 }
337 validate_with_events(&self.name, &events, "transcription", || {
338 validate_transcription_response(&data)
339 })?;
340 ProviderResponse {
341 data: InvocationData::Json(data),
342 metadata: json!({ "contentType": "application/json" }),
343 state: None,
344 }
345 }
346 };
347 if cancellation.is_cancelled() {
348 return Err(RuntimeError::Cancelled);
349 }
350 Ok(response)
351 })
352 }
353}
354
355pub async fn openai_chat_completion(
356 base_url: &str,
357 token: Option<&str>,
358 model: &str,
359 request: &Value,
360 persona: &Value,
361) -> Result<Value, String> {
362 openai_chat_completion_result(base_url, token, model, request, persona)
363 .await
364 .map_err(JsonProviderError::into_message)
365}
366
367async fn openai_chat_completion_result(
368 base_url: &str,
369 token: Option<&str>,
370 model: &str,
371 request: &Value,
372 persona: &Value,
373) -> Result<Value, JsonProviderError> {
374 let mut request = request.clone();
375 apply_persona_to_chat_request(&mut request, persona).map_err(JsonProviderError::Provider)?;
376 request
377 .as_object_mut()
378 .ok_or_else(|| {
379 JsonProviderError::Provider("chat completion request must be an object".to_string())
380 })?
381 .insert("model".to_string(), Value::String(model.to_string()));
382
383 let client = provider_http_client(None).map_err(JsonProviderError::Provider)?;
384 let response = authorized(
385 client
386 .post(provider_url(base_url, "chat/completions").map_err(JsonProviderError::Provider)?),
387 token,
388 )
389 .json(&request)
390 .send()
391 .await
392 .map_err(|error| JsonProviderError::Provider(format!("provider request failed: {error}")))?;
393 decode_json_response(response, "chat completion").await
394}
395
396pub async fn openai_embeddings(
397 base_url: &str,
398 token: Option<&str>,
399 model: &str,
400 request: &Value,
401) -> Result<Value, String> {
402 openai_embeddings_result(base_url, token, model, request)
403 .await
404 .map_err(JsonProviderError::into_message)
405}
406
407async fn openai_embeddings_result(
408 base_url: &str,
409 token: Option<&str>,
410 model: &str,
411 request: &Value,
412) -> Result<Value, JsonProviderError> {
413 let mut request = request.clone();
414 request
415 .as_object_mut()
416 .ok_or_else(|| {
417 JsonProviderError::Provider("embedding request must be an object".to_string())
418 })?
419 .insert("model".to_string(), Value::String(model.to_string()));
420
421 let client = provider_http_client(None).map_err(JsonProviderError::Provider)?;
422 let response = authorized(
423 client.post(provider_url(base_url, "embeddings").map_err(JsonProviderError::Provider)?),
424 token,
425 )
426 .json(&request)
427 .send()
428 .await
429 .map_err(|error| {
430 JsonProviderError::Provider(format!("embedding provider request failed: {error}"))
431 })?;
432 decode_json_response(response, "embedding").await
433}
434
435pub fn apply_persona_to_chat_request(request: &mut Value, persona: &Value) -> Result<(), String> {
436 let object = request
437 .as_object_mut()
438 .ok_or_else(|| "chat completion request must be an object".to_string())?;
439 apply_persona(object, persona)
440}
441
442pub async fn elevenlabs_speech(
443 base_url: &str,
444 token: Option<&str>,
445 voice_id: &str,
446 request: &Value,
447) -> Result<BinaryProviderResponse, String> {
448 let url = format!(
449 "{}/text-to-speech/{}",
450 base_url.trim_end_matches('/'),
451 encode_path_segment(voice_id)?
452 );
453 let client = provider_http_client(None)?;
454 let response = authorized(client.post(url), token)
455 .header("xi-api-key", token.unwrap_or_default())
456 .json(request)
457 .send()
458 .await
459 .map_err(|error| format!("speech provider request failed: {error}"))?;
460 decode_binary_response(response, "speech synthesis").await
461}
462
463pub async fn openai_audio_transcription(
464 base_url: &str,
465 token: Option<&str>,
466 model: &str,
467 audio: Vec<u8>,
468 file_name: &str,
469 content_type: &str,
470) -> Result<Value, String> {
471 openai_audio_transcription_result(base_url, token, model, audio, file_name, content_type)
472 .await
473 .map_err(JsonProviderError::into_message)
474}
475
476async fn openai_audio_transcription_result(
477 base_url: &str,
478 token: Option<&str>,
479 model: &str,
480 audio: Vec<u8>,
481 file_name: &str,
482 content_type: &str,
483) -> Result<Value, JsonProviderError> {
484 let part = reqwest::multipart::Part::bytes(audio)
485 .file_name(file_name.to_string())
486 .mime_str(content_type)
487 .map_err(|error| {
488 JsonProviderError::Provider(format!("audio content type is invalid: {error}"))
489 })?;
490 let form = reqwest::multipart::Form::new()
491 .text("model", model.to_string())
492 .part("file", part);
493 let client = provider_http_client(None).map_err(JsonProviderError::Provider)?;
494 let response = authorized(
495 client.post(
496 provider_url(base_url, "audio/transcriptions").map_err(JsonProviderError::Provider)?,
497 ),
498 token,
499 )
500 .multipart(form)
501 .send()
502 .await
503 .map_err(|error| {
504 JsonProviderError::Provider(format!("transcription provider request failed: {error}"))
505 })?;
506 decode_json_response(response, "audio transcription").await
507}
508
509pub async fn probe_openai_compatible(base_url: &str, token: Option<&str>) -> Result<(), String> {
510 let client = provider_http_client(Some(PROVIDER_PROBE_TIMEOUT))?;
511 let response = authorized(client.get(provider_url(base_url, "models")?), token)
512 .send()
513 .await
514 .map_err(|error| format!("provider probe failed: {error}"))?;
515 require_success(response, "probe").await
516}
517
518pub async fn probe_elevenlabs(base_url: &str, token: Option<&str>) -> Result<(), String> {
519 let client = provider_http_client(Some(PROVIDER_PROBE_TIMEOUT))?;
520 let response = authorized(client.get(provider_url(base_url, "models")?), token)
521 .header("xi-api-key", token.unwrap_or_default())
522 .send()
523 .await
524 .map_err(|error| format!("provider probe failed: {error}"))?;
525 require_success(response, "probe").await
526}
527
528#[cfg(feature = "local-whisper")]
529pub fn local_whisper_transcription(
530 model_path: &Path,
531 wav: &[u8],
532 language: Option<&str>,
533) -> Result<String, String> {
534 use std::io::Cursor;
535
536 use whisper_rs::{FullParams, SamplingStrategy, WhisperContext, WhisperContextParameters};
537
538 let mut reader = hound::WavReader::new(Cursor::new(wav))
539 .map_err(|error| format!("audio must be a valid WAV file: {error}"))?;
540 let spec = reader.spec();
541 let channels = usize::from(spec.channels);
542 if channels == 0 || spec.sample_rate == 0 {
543 return Err("WAV audio has an invalid channel count or sample rate".to_string());
544 }
545 let interleaved = match spec.sample_format {
546 hound::SampleFormat::Float => reader
547 .samples::<f32>()
548 .collect::<Result<Vec<_>, _>>()
549 .map_err(|error| format!("WAV samples could not be decoded: {error}"))?,
550 hound::SampleFormat::Int => {
551 let scale = 2_f32.powi(i32::from(spec.bits_per_sample.saturating_sub(1)));
552 reader
553 .samples::<i32>()
554 .map(|sample| {
555 sample
556 .map(|sample| sample as f32 / scale)
557 .map_err(|error| format!("WAV samples could not be decoded: {error}"))
558 })
559 .collect::<Result<Vec<_>, _>>()?
560 }
561 };
562 let mono = interleaved
563 .chunks(channels)
564 .map(|frame| frame.iter().copied().sum::<f32>() / frame.len() as f32)
565 .collect::<Vec<_>>();
566 let samples = resample_linear(&mono, spec.sample_rate, 16_000);
567 if samples.is_empty() {
568 return Err("WAV audio does not contain samples".to_string());
569 }
570
571 let model_path = model_path
572 .to_str()
573 .ok_or_else(|| "Whisper model path is not valid UTF-8".to_string())?;
574 let context = WhisperContext::new_with_params(model_path, WhisperContextParameters::default())
575 .map_err(|error| format!("Whisper model could not be loaded: {error}"))?;
576 let mut state = context
577 .create_state()
578 .map_err(|error| format!("Whisper state could not be created: {error}"))?;
579 let mut params = FullParams::new(SamplingStrategy::Greedy { best_of: 1 });
580 params.set_print_progress(false);
581 params.set_print_realtime(false);
582 params.set_print_timestamps(false);
583 params.set_language(language);
584 state
585 .full(params, &samples)
586 .map_err(|error| format!("Whisper transcription failed: {error}"))?;
587 let segments = state
588 .as_iter()
589 .map(|segment| {
590 segment
591 .to_str_lossy()
592 .map(|text| text.into_owned())
593 .map_err(|error| format!("Whisper segment could not be decoded: {error}"))
594 })
595 .collect::<Result<Vec<_>, _>>()?;
596 Ok(segments.join("").trim().to_string())
597}
598
599#[cfg(not(feature = "local-whisper"))]
600pub fn local_whisper_transcription(
601 _model_path: &Path,
602 _wav: &[u8],
603 _language: Option<&str>,
604) -> Result<String, String> {
605 Err("this Vifu build does not include local Whisper support".to_string())
606}
607
608pub fn resolve_local_model_path(home_dir: &Path, model: &str) -> Result<PathBuf, String> {
609 let model = model.trim();
610 if model.is_empty()
611 || model.len() > 255
612 || model.contains('/')
613 || model.contains('\\')
614 || model == "."
615 || model == ".."
616 {
617 return Err("local model must be a file name inside ~/.vifu/models".to_string());
618 }
619 Ok(home_dir.join("models").join(model))
620}
621
622fn apply_persona(
623 request: &mut serde_json::Map<String, Value>,
624 persona: &Value,
625) -> Result<(), String> {
626 let prompt = persona_prompt(persona);
627 if prompt.is_empty() {
628 return Ok(());
629 }
630 let messages = request
631 .get_mut("messages")
632 .and_then(Value::as_array_mut)
633 .ok_or_else(|| "chat completion messages must be an array".to_string())?;
634 messages.insert(0, json!({ "role": "system", "content": prompt }));
635 Ok(())
636}
637
638fn persona_prompt(persona: &Value) -> String {
639 let mut sections = Vec::new();
640 if let Some(prompt) = persona
641 .get("systemPrompt")
642 .and_then(Value::as_str)
643 .map(str::trim)
644 .filter(|value| !value.is_empty())
645 {
646 sections.push(prompt.to_string());
647 }
648 if let Some(files) = persona.get("files").and_then(Value::as_object) {
649 for (name, content) in files {
650 let Some(content) = content
651 .as_str()
652 .map(str::trim)
653 .filter(|value| !value.is_empty())
654 else {
655 continue;
656 };
657 sections.push(format!("# {name}\n\n{content}"));
658 }
659 }
660 sections.join("\n\n")
661}
662
663enum JsonProviderError {
664 Provider(String),
665 MalformedResponse(String),
666}
667
668impl JsonProviderError {
669 fn into_message(self) -> String {
670 match self {
671 Self::Provider(message) | Self::MalformedResponse(message) => message,
672 }
673 }
674}
675
676fn provider_json_result(
677 provider_name: &str,
678 events: &ProviderEventSink,
679 kind: &str,
680 result: Result<Value, JsonProviderError>,
681) -> Result<Value, RuntimeError> {
682 match result {
683 Ok(response) => Ok(response),
684 Err(JsonProviderError::Provider(message)) => {
685 Err(RuntimeError::provider(provider_name, message))
686 }
687 Err(JsonProviderError::MalformedResponse(message)) => {
688 let started = Instant::now();
689 events.stage_started(ProviderStage::Validate, json!({ "kind": kind }));
690 Err(validation_error(
691 provider_name,
692 events,
693 kind,
694 started,
695 message,
696 ))
697 }
698 }
699}
700
701fn validate_with_events(
702 provider_name: &str,
703 events: &ProviderEventSink,
704 kind: &str,
705 validate: impl FnOnce() -> Result<Value, String>,
706) -> Result<(), RuntimeError> {
707 let started = Instant::now();
708 events.stage_started(ProviderStage::Validate, json!({ "kind": kind }));
709 match validate() {
710 Ok(metadata) => {
711 events.stage_completed(ProviderStage::Validate, elapsed_ms(started), metadata);
712 Ok(())
713 }
714 Err(message) => Err(validation_error(
715 provider_name,
716 events,
717 kind,
718 started,
719 message,
720 )),
721 }
722}
723
724fn validation_error(
725 provider_name: &str,
726 events: &ProviderEventSink,
727 kind: &str,
728 started: Instant,
729 message: String,
730) -> RuntimeError {
731 let error = RuntimeError::provider(provider_name, message);
732 events.stage_failed(
733 ProviderStage::Validate,
734 elapsed_ms(started),
735 error.to_string(),
736 json!({ "kind": kind }),
737 );
738 error
739}
740
741fn validate_openai_chat_response(response: &Value) -> Result<Value, String> {
742 let choices = response
743 .get("choices")
744 .and_then(Value::as_array)
745 .filter(|choices| !choices.is_empty())
746 .ok_or_else(|| "chat response has no choices".to_string())?;
747 let message = choices[0]
748 .get("message")
749 .and_then(Value::as_object)
750 .ok_or_else(|| "chat response first choice has no assistant message".to_string())?;
751 if let Some(role) = message.get("role") {
752 if role.as_str() != Some("assistant") {
753 return Err("chat response first message is not from the assistant".to_string());
754 }
755 }
756 let content = message.get("content");
757 if let Some(content) = content {
758 if !matches!(content, Value::Null | Value::String(_) | Value::Array(_)) {
759 return Err("chat response assistant content has an invalid type".to_string());
760 }
761 }
762 let tool_calls = message.get("tool_calls");
763 if tool_calls.is_some_and(|calls| !calls.is_array()) {
764 return Err("chat response assistant tool_calls is not an array".to_string());
765 }
766 let function_call = message.get("function_call");
767 if function_call.is_some_and(|call| !call.is_object() && !call.is_null()) {
768 return Err("chat response assistant function_call is not an object".to_string());
769 }
770 if content.is_none() && tool_calls.is_none() && function_call.is_none() {
771 return Err("chat response assistant message has no content or tool calls".to_string());
772 }
773 Ok(json!({
774 "kind": "chat",
775 "choices": choices.len(),
776 "toolCalls": tool_calls.and_then(Value::as_array).map_or(0, Vec::len),
777 }))
778}
779
780fn validate_openai_embedding_response(response: &Value) -> Result<Value, String> {
781 let rows = response
782 .get("data")
783 .and_then(Value::as_array)
784 .filter(|rows| !rows.is_empty())
785 .ok_or_else(|| "embedding response has no data rows".to_string())?;
786 let mut dimensions = None;
787 let mut encoding = None;
788 for (index, row) in rows.iter().enumerate() {
789 let embedding = row
790 .get("embedding")
791 .ok_or_else(|| format!("embedding row {index} has no vector"))?;
792 let (row_dimensions, row_encoding) = if let Some(values) = embedding.as_array() {
793 if values.is_empty()
794 || values
795 .iter()
796 .any(|value| !value.as_f64().is_some_and(f64::is_finite))
797 {
798 return Err(format!(
799 "embedding row {index} has an invalid numeric vector"
800 ));
801 }
802 (values.len(), "float")
803 } else if let Some(encoded) = embedding.as_str() {
804 let row_dimensions = decode_base64_float32_dimensions(encoded)
805 .map_err(|message| format!("embedding row {index} {message}"))?;
806 (row_dimensions, "base64")
807 } else {
808 return Err(format!("embedding row {index} has an invalid vector type"));
809 };
810 if let Some(expected_dimensions) = dimensions {
811 if expected_dimensions != row_dimensions {
812 return Err(format!(
813 "embedding row {index} changed dimension from {expected_dimensions} to {row_dimensions}"
814 ));
815 }
816 } else {
817 dimensions = Some(row_dimensions);
818 }
819 if let Some(expected_encoding) = encoding {
820 if expected_encoding != row_encoding {
821 return Err(format!(
822 "embedding row {index} changed encoding from {expected_encoding} to {row_encoding}"
823 ));
824 }
825 } else {
826 encoding = Some(row_encoding);
827 }
828 }
829 Ok(json!({
830 "kind": "embedding",
831 "rows": rows.len(),
832 "dimensions": dimensions,
833 "encoding": encoding,
834 }))
835}
836
837fn validate_transcription_response(response: &Value) -> Result<Value, String> {
838 let text = response
839 .get("text")
840 .and_then(Value::as_str)
841 .ok_or_else(|| "transcription response text is missing or is not a string".to_string())?;
842 Ok(json!({ "kind": "transcription", "characters": text.chars().count() }))
843}
844
845fn decode_base64_float32_dimensions(encoded: &str) -> Result<usize, String> {
846 let mut bytes = encoded.as_bytes().to_vec();
847 if bytes.is_empty() || bytes.len() % 4 == 1 {
848 return Err("has invalid base64 float32 data".to_string());
849 }
850 match bytes.len() % 4 {
851 2 => bytes.extend_from_slice(b"=="),
852 3 => bytes.push(b'='),
853 _ => {}
854 }
855 let mut decoded = Vec::with_capacity(bytes.len() / 4 * 3);
856 let chunks = bytes.len() / 4;
857 for (chunk_index, chunk) in bytes.chunks_exact(4).enumerate() {
858 let last = chunk_index + 1 == chunks;
859 let padding = chunk.iter().rev().take_while(|byte| **byte == b'=').count();
860 if padding > 2 || (!last && padding > 0) || chunk[..4 - padding].contains(&b'=') {
861 return Err("has invalid base64 padding".to_string());
862 }
863 let a = base64_value(chunk[0])?;
864 let b = base64_value(chunk[1])?;
865 let c = if padding >= 2 {
866 0
867 } else {
868 base64_value(chunk[2])?
869 };
870 let d = if padding >= 1 {
871 0
872 } else {
873 base64_value(chunk[3])?
874 };
875 if (padding == 2 && b & 0x0f != 0) || (padding == 1 && c & 0x03 != 0) {
876 return Err("has non-canonical base64 padding".to_string());
877 }
878 decoded.push((a << 2) | (b >> 4));
879 if padding < 2 {
880 decoded.push((b << 4) | (c >> 2));
881 }
882 if padding == 0 {
883 decoded.push((c << 6) | d);
884 }
885 }
886 if decoded.is_empty() || decoded.len() % std::mem::size_of::<f32>() != 0 {
887 return Err("does not contain whole float32 values".to_string());
888 }
889 for bytes in decoded.chunks_exact(4) {
890 let value = f32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]);
891 if !value.is_finite() {
892 return Err("contains a non-finite float32 value".to_string());
893 }
894 }
895 Ok(decoded.len() / std::mem::size_of::<f32>())
896}
897
898fn base64_value(byte: u8) -> Result<u8, String> {
899 match byte {
900 b'A'..=b'Z' => Ok(byte - b'A'),
901 b'a'..=b'z' => Ok(byte - b'a' + 26),
902 b'0'..=b'9' => Ok(byte - b'0' + 52),
903 b'+' | b'-' => Ok(62),
904 b'/' | b'_' => Ok(63),
905 _ => Err("has invalid base64 characters".to_string()),
906 }
907}
908
909fn elapsed_ms(started: Instant) -> u64 {
910 u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX)
911}
912
913fn provider_url(base_url: &str, path: &str) -> Result<String, String> {
914 let base_url = base_url.trim();
915 if !(base_url.starts_with("http://") || base_url.starts_with("https://")) {
916 return Err("provider URL must use http or https".to_string());
917 }
918 Ok(format!("{}/{}", base_url.trim_end_matches('/'), path))
919}
920
921fn provider_http_client(timeout: Option<std::time::Duration>) -> Result<reqwest::Client, String> {
922 let mut builder = reqwest::Client::builder().redirect(reqwest::redirect::Policy::none());
923 if let Some(timeout) = timeout {
924 builder = builder.timeout(timeout);
925 }
926 builder
927 .build()
928 .map_err(|error| format!("provider client could not be created: {error}"))
929}
930
931fn authorized(builder: reqwest::RequestBuilder, token: Option<&str>) -> reqwest::RequestBuilder {
932 match token.map(str::trim).filter(|token| !token.is_empty()) {
933 Some(token) => builder.bearer_auth(token),
934 None => builder,
935 }
936}
937
938async fn decode_json_response(
939 response: reqwest::Response,
940 operation: &str,
941) -> Result<Value, JsonProviderError> {
942 let status = response.status();
943 let body = response.bytes().await.map_err(|error| {
944 JsonProviderError::Provider(format!("{operation} response could not be read: {error}"))
945 })?;
946 if !status.is_success() {
947 return Err(JsonProviderError::Provider(provider_error(
948 operation,
949 status.as_u16(),
950 &body,
951 )));
952 }
953 serde_json::from_slice(&body).map_err(|error| {
954 JsonProviderError::MalformedResponse(format!(
955 "{operation} response is not valid JSON: {error}"
956 ))
957 })
958}
959
960async fn decode_binary_response(
961 response: reqwest::Response,
962 operation: &str,
963) -> Result<BinaryProviderResponse, String> {
964 let status = response.status();
965 let content_type = response_content_type(response.headers());
966 let body = response
967 .bytes()
968 .await
969 .map_err(|error| format!("{operation} response could not be read: {error}"))?;
970 if !status.is_success() {
971 return Err(provider_error(operation, status.as_u16(), &body));
972 }
973 Ok(BinaryProviderResponse {
974 content_type,
975 body: body.to_vec(),
976 })
977}
978
979async fn require_success(response: reqwest::Response, operation: &str) -> Result<(), String> {
980 let status = response.status();
981 if status.is_success() {
982 return Ok(());
983 }
984 let body = response
985 .bytes()
986 .await
987 .map_err(|error| format!("provider {operation} response could not be read: {error}"))?;
988 Err(provider_error(operation, status.as_u16(), &body))
989}
990
991fn response_content_type(headers: &HeaderMap) -> String {
992 headers
993 .get(CONTENT_TYPE)
994 .and_then(|value| value.to_str().ok())
995 .unwrap_or("application/octet-stream")
996 .to_string()
997}
998
999fn provider_error(operation: &str, status: u16, body: &[u8]) -> String {
1000 let message = serde_json::from_slice::<Value>(body)
1001 .ok()
1002 .and_then(|value| {
1003 value
1004 .pointer("/error/message")
1005 .or_else(|| value.get("error"))
1006 .and_then(Value::as_str)
1007 .map(str::trim)
1008 .filter(|value| !value.is_empty())
1009 .map(|value| value.chars().take(512).collect::<String>())
1010 })
1011 .unwrap_or_else(|| format!("HTTP {status}"));
1012 format!("provider {operation} failed: {message}")
1013}
1014
1015fn encode_path_segment(value: &str) -> Result<String, String> {
1016 let value = value.trim();
1017 if value.is_empty()
1018 || value.len() > 256
1019 || !value
1020 .bytes()
1021 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_'))
1022 {
1023 return Err("provider resource ID contains unsupported characters".to_string());
1024 }
1025 Ok(value.to_string())
1026}
1027
1028#[cfg(feature = "local-whisper")]
1029fn resample_linear(input: &[f32], from_hz: u32, to_hz: u32) -> Vec<f32> {
1030 if input.is_empty() || from_hz == 0 || to_hz == 0 {
1031 return Vec::new();
1032 }
1033 if from_hz == to_hz {
1034 return input.to_vec();
1035 }
1036 let output_len = (input.len() as u64 * u64::from(to_hz) / u64::from(from_hz)) as usize;
1037 (0..output_len)
1038 .map(|index| {
1039 let source = index as f64 * f64::from(from_hz) / f64::from(to_hz);
1040 let left = source.floor() as usize;
1041 let right = (left + 1).min(input.len() - 1);
1042 let fraction = (source - left as f64) as f32;
1043 input[left] + (input[right] - input[left]) * fraction
1044 })
1045 .collect()
1046}
1047
1048#[cfg(test)]
1049mod tests {
1050 use std::io::{Read, Write};
1051 use std::net::TcpListener;
1052 use std::sync::{Arc, Mutex};
1053 use std::thread;
1054
1055 use serde_json::json;
1056
1057 use super::{
1058 openai_chat_completion, persona_prompt, probe_openai_compatible, provider_json_result,
1059 provider_url, resolve_local_model_path, validate_openai_chat_response,
1060 validate_openai_embedding_response, validate_transcription_response, validate_with_events,
1061 JsonProviderError,
1062 };
1063 use crate::{ProviderEvent, ProviderEventSink, ProviderStage, RuntimeError};
1064
1065 #[test]
1066 fn builds_a_portable_persona_prompt() {
1067 assert_eq!(
1068 persona_prompt(&json!({
1069 "systemPrompt": "Stay concise.",
1070 "files": { "SOUL.md": "You are the steward." }
1071 })),
1072 "Stay concise.\n\n# SOUL.md\n\nYou are the steward."
1073 );
1074 }
1075
1076 #[test]
1077 fn appends_openai_compatible_paths() {
1078 assert_eq!(
1079 provider_url("https://example.com/v1/", "chat/completions").unwrap(),
1080 "https://example.com/v1/chat/completions"
1081 );
1082 }
1083
1084 #[test]
1085 fn appends_the_openai_embedding_path() {
1086 assert_eq!(
1087 provider_url("https://example.com/v1/", "embeddings").unwrap(),
1088 "https://example.com/v1/embeddings"
1089 );
1090 }
1091
1092 #[tokio::test]
1093 async fn provider_requests_do_not_follow_redirects() {
1094 let (base_url, server) = redirecting_provider();
1095
1096 let error = openai_chat_completion(
1097 &base_url,
1098 Some("test-token"),
1099 "local-model",
1100 &json!({"messages": [{"role": "user", "content": "hello"}]}),
1101 &json!({}),
1102 )
1103 .await
1104 .unwrap_err();
1105
1106 server.join().unwrap();
1107 assert!(error.contains("307"), "unexpected provider error: {error}");
1108 }
1109
1110 #[tokio::test]
1111 async fn provider_probes_do_not_follow_redirects() {
1112 let (base_url, server) = redirecting_provider();
1113
1114 let error = probe_openai_compatible(&base_url, Some("test-token"))
1115 .await
1116 .unwrap_err();
1117
1118 server.join().unwrap();
1119 assert!(error.contains("307"), "unexpected probe error: {error}");
1120 }
1121
1122 fn redirecting_provider() -> (String, thread::JoinHandle<()>) {
1123 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1124 let address = listener.local_addr().unwrap();
1125 let server = thread::spawn(move || {
1126 let (mut stream, _) = listener.accept().unwrap();
1127 let mut request = [0_u8; 4096];
1128 let _ = stream.read(&mut request);
1129 stream
1130 .write_all(
1131 b"HTTP/1.1 307 Temporary Redirect\r\nLocation: http://127.0.0.1:9/exfil\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
1132 )
1133 .unwrap();
1134 });
1135 (format!("http://{address}/v1"), server)
1136 }
1137
1138 #[test]
1139 fn keeps_local_models_inside_the_vifu_model_directory() {
1140 let path =
1141 resolve_local_model_path(std::path::Path::new("/tmp/.vifu"), "tiny.bin").unwrap();
1142 assert_eq!(path, std::path::Path::new("/tmp/.vifu/models/tiny.bin"));
1143 assert!(resolve_local_model_path(std::path::Path::new("/tmp/.vifu"), "../key").is_err());
1144 }
1145
1146 #[test]
1147 fn validates_openai_chat_assistant_text() {
1148 let metadata = validate_openai_chat_response(&json!({
1149 "choices": [{
1150 "message": {
1151 "role": "assistant",
1152 "content": [{ "type": "text", "text": "Ready." }]
1153 }
1154 }]
1155 }))
1156 .unwrap();
1157
1158 assert_eq!(
1159 metadata,
1160 json!({ "kind": "chat", "choices": 1, "toolCalls": 0 })
1161 );
1162 }
1163
1164 #[test]
1165 fn accepts_empty_and_tool_call_only_chat_outputs() {
1166 let empty = validate_openai_chat_response(&json!({
1167 "choices": [{ "message": { "role": "assistant", "content": "" } }]
1168 }))
1169 .unwrap();
1170 let tool_call = validate_openai_chat_response(&json!({
1171 "choices": [{
1172 "message": {
1173 "role": "assistant",
1174 "content": null,
1175 "tool_calls": [{
1176 "id": "call-1",
1177 "type": "function",
1178 "function": { "name": "move", "arguments": "{}" }
1179 }]
1180 }
1181 }]
1182 }))
1183 .unwrap();
1184
1185 assert_eq!(empty["toolCalls"], 0);
1186 assert_eq!(tool_call["toolCalls"], 1);
1187 }
1188
1189 #[test]
1190 fn validates_consistent_numeric_embedding_rows() {
1191 let metadata = validate_openai_embedding_response(&json!({
1192 "data": [
1193 { "embedding": [1, 2.5] },
1194 { "embedding": [-3, 4] }
1195 ]
1196 }))
1197 .unwrap();
1198
1199 assert_eq!(
1200 metadata,
1201 json!({
1202 "kind": "embedding",
1203 "rows": 2,
1204 "dimensions": 2,
1205 "encoding": "float"
1206 })
1207 );
1208 }
1209
1210 #[test]
1211 fn validates_consistent_base64_float32_embedding_rows() {
1212 let metadata = validate_openai_embedding_response(&json!({
1213 "data": [
1214 { "embedding": "AACAPwAAAMA=" },
1215 { "embedding": "AACAPwAAAMA=" }
1216 ]
1217 }))
1218 .unwrap();
1219
1220 assert_eq!(metadata["dimensions"], 2);
1221 assert_eq!(metadata["encoding"], "base64");
1222 }
1223
1224 #[test]
1225 fn rejects_embedding_rows_with_different_dimensions() {
1226 let error = validate_openai_embedding_response(&json!({
1227 "data": [
1228 { "embedding": [1, 2] },
1229 { "embedding": [3] }
1230 ]
1231 }))
1232 .unwrap_err();
1233
1234 assert_eq!(error, "embedding row 1 changed dimension from 2 to 1");
1235 }
1236
1237 #[test]
1238 fn rejects_malformed_base64_embedding_vectors() {
1239 let error = validate_openai_embedding_response(&json!({
1240 "data": [{ "embedding": "not base64" }]
1241 }))
1242 .unwrap_err();
1243
1244 assert!(error.contains("invalid base64"));
1245 }
1246
1247 #[test]
1248 fn accepts_silence_and_rejects_transcription_without_a_text_field() {
1249 assert_eq!(
1250 validate_transcription_response(&json!({ "text": "" })).unwrap(),
1251 json!({ "kind": "transcription", "characters": 0 })
1252 );
1253 let error = validate_transcription_response(&json!({})).unwrap_err();
1254
1255 assert_eq!(
1256 error,
1257 "transcription response text is missing or is not a string"
1258 );
1259 }
1260
1261 #[test]
1262 fn malformed_chat_emits_validate_failed_and_returns_provider_error() {
1263 let captured = Arc::new(Mutex::new(Vec::new()));
1264 let event_capture = Arc::clone(&captured);
1265 let events = ProviderEventSink::from_fn(move |event| {
1266 event_capture.lock().unwrap().push(event);
1267 });
1268
1269 let error = validate_with_events("remote", &events, "chat", || {
1270 validate_openai_chat_response(&json!({
1271 "choices": [{ "message": { "role": "assistant", "content": 42 } }]
1272 }))
1273 })
1274 .unwrap_err();
1275
1276 assert!(matches!(
1277 error,
1278 RuntimeError::Provider { ref provider, ref message }
1279 if provider == "remote"
1280 && message == "chat response assistant content has an invalid type"
1281 ));
1282 let captured = captured.lock().unwrap();
1283 assert_eq!(captured.len(), 2);
1284 assert!(matches!(
1285 captured[0],
1286 ProviderEvent::StageStarted {
1287 stage: ProviderStage::Validate,
1288 ..
1289 }
1290 ));
1291 assert!(matches!(
1292 captured[1],
1293 ProviderEvent::StageFailed {
1294 stage: ProviderStage::Validate,
1295 ..
1296 }
1297 ));
1298 }
1299
1300 #[test]
1301 fn malformed_embedding_emits_validate_failed_and_returns_provider_error() {
1302 let captured = Arc::new(Mutex::new(Vec::new()));
1303 let event_capture = Arc::clone(&captured);
1304 let events = ProviderEventSink::from_fn(move |event| {
1305 event_capture.lock().unwrap().push(event);
1306 });
1307
1308 let error = validate_with_events("remote", &events, "embedding", || {
1309 validate_openai_embedding_response(&json!({
1310 "data": [
1311 { "embedding": [1, 2] },
1312 { "embedding": [3] }
1313 ]
1314 }))
1315 })
1316 .unwrap_err();
1317
1318 assert!(matches!(
1319 error,
1320 RuntimeError::Provider { ref provider, ref message }
1321 if provider == "remote"
1322 && message == "embedding row 1 changed dimension from 2 to 1"
1323 ));
1324 let captured = captured.lock().unwrap();
1325 assert_eq!(captured.len(), 2);
1326 assert!(matches!(
1327 captured[0],
1328 ProviderEvent::StageStarted {
1329 stage: ProviderStage::Validate,
1330 ..
1331 }
1332 ));
1333 assert!(matches!(
1334 captured[1],
1335 ProviderEvent::StageFailed {
1336 stage: ProviderStage::Validate,
1337 ..
1338 }
1339 ));
1340 }
1341
1342 #[test]
1343 fn invalid_json_output_emits_validate_failed() {
1344 let captured = Arc::new(Mutex::new(Vec::new()));
1345 let event_capture = Arc::clone(&captured);
1346 let events = ProviderEventSink::from_fn(move |event| {
1347 event_capture.lock().unwrap().push(event);
1348 });
1349
1350 let error = provider_json_result(
1351 "remote",
1352 &events,
1353 "chat",
1354 Err(JsonProviderError::MalformedResponse(
1355 "chat completion response is not valid JSON".to_string(),
1356 )),
1357 )
1358 .unwrap_err();
1359
1360 assert!(matches!(error, RuntimeError::Provider { .. }));
1361 let captured = captured.lock().unwrap();
1362 assert!(matches!(
1363 captured.as_slice(),
1364 [
1365 ProviderEvent::StageStarted {
1366 stage: ProviderStage::Validate,
1367 ..
1368 },
1369 ProviderEvent::StageFailed {
1370 stage: ProviderStage::Validate,
1371 ..
1372 }
1373 ]
1374 ));
1375 }
1376}