Skip to main content

vifu_runtime/
providers.rs

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/// Capability protocol used by [`HttpCapabilityProvider`].
22#[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
82/// A runtime-registered provider assembled from capability protocol routes.
83///
84/// This is one provider object regardless of vendor count. Add routes at
85/// runtime instead of selecting provider-specific Cargo features.
86pub 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}