Skip to main content

rig_core/providers/openai/wire/
modality.rs

1//! Embedding, transcription, image, speech, listing, reranking, and verification wires.
2//! Each operation consumes a whole reply, using its dialect's JSON or binary format.
3//!
4//! ```
5//! use rig_core::providers::openai::{OpenAI, TEXT_EMBEDDING_3_SMALL};
6//! let wire = OpenAI::new("key").embedding(TEXT_EMBEDDING_3_SMALL, None);
7//! ```
8
9use crate::wire::Flow;
10use serde::{Deserialize, Serialize};
11
12use crate::embeddings;
13use crate::error::EncodeError;
14use crate::error::ProviderError;
15use crate::model::{ModelInfo, ModelList};
16use crate::operation::{
17    Embedding, ModelListing, ModelPage, Rerank as RerankOp, Transcription, Verify as VerifyOp,
18};
19use crate::providers::internal::wire::classify_untyped_line;
20use crate::providers::openai::embedding::Usage;
21use crate::providers::openai::embedding::{
22    CompatibleEmbeddingResponse, EncodingFormat, model_dimensions_from_identifier,
23};
24use crate::transcription::TranscriptionRequest;
25use crate::wire::{
26    Body, Capabilities, Decoder, Descriptor, Encoded, Framing, Mode, Out, Wire, WireEvent,
27    WireFrame,
28};
29
30use super::OpenAIConfig;
31
32/// Paired request and response formats for image generation.
33#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
34pub enum ImageBody {
35    /// `{model, prompt, size}`, answered with `data[].b64_json`.
36    #[default]
37    OpenAi,
38    /// xAI: `{model, prompt, response_format, aspect_ratio}` and no `size`,
39    /// answered with `data[].b64_json` and no `created`.
40    Xai,
41    /// `{model_name, prompt, height, width}`, answered with `images[].image`.
42    Hyperbolic,
43    /// `{model, prompt, width, height}`, answered with base64 strings in `images`.
44    Venice,
45    /// `{inputs, parameters: {width, height}}`, answered with raw image bytes.
46    /// The model is addressed through the URL path.
47    HuggingFace,
48}
49
50/// Which body a speech endpoint takes.
51#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
52pub enum SpeechBody {
53    /// OpenAI: `{model, input, voice, speed}`.
54    #[default]
55    OpenAi,
56    /// xAI: `{text, voice_id, language}`, with `eve` as the default voice.
57    Xai,
58    /// Hyperbolic: `{language, speaker, text, speed}`, answered with
59    /// `{"audio": "<base64>"}` rather than the audio bytes themselves.
60    ///
61    /// It addresses this endpoint by *language*, so the identifier a caller
62    /// passes as the model is the language tag (`"EN"`).
63    Hyperbolic,
64}
65
66/// Which body a transcription endpoint takes.
67#[derive(Clone, Copy, Debug, PartialEq, Eq)]
68pub enum TranscriptionBody {
69    /// OpenAI: a `multipart/form-data` upload with the audio as a file part
70    /// beside `model`, `language`, `prompt` and `temperature`.
71    Multipart,
72    /// OpenRouter: a JSON body whose audio rides base64-encoded under
73    /// `input_audio`, with its container format beside it
74    /// (`{"input_audio": {"data": "…", "format": "mp3"}, "model": …}`).
75    /// The gateway's speech-to-text route serves only this shape, and has no
76    /// top-level `prompt` field at all.
77    InputAudioJson,
78}
79
80/// Which field a dialect takes an embedding width in.
81#[derive(Clone, Copy, Debug, PartialEq, Eq)]
82pub enum DimensionsField {
83    /// The OpenAI-compatible `dimensions` field.
84    Dimensions,
85    /// Mistral's `output_dimension`.
86    OutputDimension,
87    /// The server ignores any width field, so none is sent (`llama-server`
88    /// reads no such field and would answer 200 with the native width).
89    Ignored,
90}
91
92impl DimensionsField {
93    /// The body field a requested width goes in, or `None` when the dialect
94    /// reads no width field at all.
95    ///
96    /// The encoder puts a width in this field and a refusal names it, so
97    /// both spell it from here rather than from two matching literals.
98    pub const fn name(self) -> Option<&'static str> {
99        match self {
100            Self::Dimensions => Some("dimensions"),
101            Self::OutputDimension => Some("output_dimension"),
102            Self::Ignored => None,
103        }
104    }
105}
106
107/// Which widths a request may name for one embedding model.
108#[derive(Clone, Copy, Debug, PartialEq, Eq)]
109pub enum AcceptedWidths {
110    /// The model emits one width and reads no width field, so any value but
111    /// its own native width is a request for a parameter the provider does
112    /// not accept there, and is refused as a request error.
113    Fixed,
114    /// The model truncates to any width in `min..=max`, and anything else is
115    /// refused with [`requirement`](Self::Range::requirement).
116    Range {
117        /// Narrowest width the provider honours.
118        min: usize,
119        /// Widest width the provider honours.
120        max: usize,
121        /// Static error clause describing the accepted bounds.
122        /// Must agree with `min` and `max`.
123        requirement: &'static str,
124    },
125}
126
127/// One embedding model's width contract: the width it returns unasked, and
128/// the widths it will honour when asked.
129///
130/// Stated per model rather than per dialect because a dialect serves models
131/// of different widths, and a model's default is not always its maximum.
132#[derive(Clone, Copy, Debug, PartialEq, Eq)]
133#[non_exhaustive]
134pub struct ModelWidth {
135    /// The model identifier, as the `model` field spells it.
136    pub model: &'static str,
137    /// Default width reported when no width is requested, or `None` if unknown.
138    /// Unknown widths report zero as the embedding model's `ndims`.
139    pub default: Option<usize>,
140    /// The widths a request may name.
141    pub accepted: AcceptedWidths,
142}
143
144/// Reranking endpoint policy. An empty [`Self::path`] disables reranking.
145#[derive(Clone, Copy, Debug, PartialEq, Eq)]
146#[non_exhaustive]
147pub struct RerankQuirks {
148    /// The rerank path, or empty when the dialect offers none.
149    pub path: &'static str,
150    /// Most documents the provider accepts in one request.
151    pub max_documents: usize,
152    /// Whether the model is a body field.
153    pub sends_model_field: bool,
154}
155
156impl RerankQuirks {
157    /// The signal for a dialect with no reranking endpoint.
158    pub const fn unsupported() -> Self {
159        Self {
160            path: "",
161            max_documents: 0,
162            sends_model_field: true,
163        }
164    }
165}
166
167/// What a dialect's embeddings endpoint accepts.
168#[derive(Clone, Copy, Debug, PartialEq, Eq)]
169#[non_exhaustive]
170pub struct EmbeddingQuirks {
171    /// Most inputs the provider embeds in one request.
172    pub max_documents: usize,
173    /// Whether a successful reply must carry usage.
174    pub requires_usage: bool,
175    /// Whether the provider accepts `encoding_format`.
176    pub supports_encoding_format: bool,
177    /// Whether the provider accepts `user`.
178    pub supports_user: bool,
179    /// Whether the model is a body field (false for Azure, which addresses a
180    /// deployment through the URL).
181    pub sends_model_field: bool,
182    /// Which field a requested width goes in.
183    pub dimensions: DimensionsField,
184    /// Model width contracts used for capability reporting and request validation.
185    /// Consulted before the shared OpenAI model-width table.
186    pub widths: &'static [ModelWidth],
187    /// The `requirement` clause refusing a declared width of zero, or
188    /// `None` for a dialect that lets zero through as rig's own "unknown"
189    /// sentinel rather than a claim.
190    pub refuse_zero_width: Option<&'static str>,
191}
192
193impl EmbeddingQuirks {
194    /// OpenAI's own embeddings contract, which most dialects inherit.
195    pub const fn openai() -> Self {
196        Self {
197            max_documents: 1024,
198            requires_usage: true,
199            supports_encoding_format: true,
200            supports_user: true,
201            sends_model_field: true,
202            dimensions: DimensionsField::Dimensions,
203            // Shared OpenAI model widths are resolved separately.
204            widths: &[],
205            // Zero represents an unknown width, not a requested dimension.
206            refuse_zero_width: None,
207        }
208    }
209}
210
211impl super::SubRoute {
212    /// Whether this sub-provider serves the endpoints that address the model
213    /// through the URL (transcription, image generation).
214    pub fn serves_model_routed_endpoints(&self) -> bool {
215        matches!(self, Self::HFInference)
216    }
217}
218
219impl OpenAIConfig {
220    /// Set the `api-version` Azure's speech endpoint is versioned by.
221    pub fn with_audio_api_version(mut self, api_version: impl Into<String>) -> Self {
222        self.audio_api_version = Some(api_version.into());
223        self
224    }
225
226    /// The embeddings wire for `model`.
227    pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
228        Embeddings::new(self.clone(), model, ndims)
229    }
230
231    /// The rerank wire for `model`.
232    pub(crate) fn rerank(&self, model: impl Into<String>) -> Rerank {
233        Rerank::new(self.clone(), model)
234    }
235
236    /// The transcription wire for `model`.
237    pub(crate) fn transcription(&self, model: impl Into<String>) -> Transcriptions {
238        Transcriptions::new(self.clone(), model)
239    }
240
241    /// The model-listing wire.
242    pub(crate) fn models(&self) -> Models {
243        Models::new(self.clone())
244    }
245
246    /// The credential-check wire.
247    pub(crate) fn verify(&self) -> Verify {
248        Verify::new(self.clone())
249    }
250
251    /// The image-generation wire for `model`.
252    #[cfg(feature = "image")]
253    pub(crate) fn image_generation(&self, model: impl Into<String>) -> Images {
254        Images::new(self.clone(), model)
255    }
256
257    /// The speech wire for `model`.
258    #[cfg(feature = "audio")]
259    pub(crate) fn audio_generation(&self, model: impl Into<String>) -> Speech {
260        Speech::new(self.clone(), model)
261    }
262
263    /// The `api-version` a speech request carries.
264    #[cfg(feature = "audio")]
265    pub(crate) fn speech_api_version(&self) -> Option<&str> {
266        self.audio_api_version
267            .as_deref()
268            .or(self.api_version.as_deref())
269    }
270
271    /// Resolve a fixed or model-addressed modality URL.
272    /// Return an error if the selected sub-route does not serve model-routed endpoints.
273    pub(crate) fn modality_uri(
274        &self,
275        endpoint: &str,
276        fixed: &'static str,
277        model: &str,
278    ) -> Result<String, String> {
279        if !self.dialect.quirks.model_is_modality_path {
280            return Ok(self.uri(fixed, self.deployment(model)));
281        }
282        let route = self.route();
283        if !route.serves_model_routed_endpoints() {
284            return Err(format!(
285                "{endpoint} endpoint is not supported yet for {route}"
286            ));
287        }
288        Ok(format!(
289            "{}/{}",
290            self.base_url.trim_end_matches('/'),
291            model.trim_start_matches('/')
292        ))
293    }
294}
295
296/// The width contract of Mistral's embedding models.
297///
298/// `mistral-embed` is fixed at 1024 and reads no width field: Mistral
299/// answers any other value with an error rather than truncating, so a
300/// request naming one is refused before it is built. Codestral Embed is
301/// configurable up to 3072 and takes its width as `output_dimension`.
302///
303/// The dated aliases are listed beside their rolling names because a caller
304/// pinning `mistral-embed-2312` gets the same model, and a model absent from
305/// this table reports `ndims() == 0`.
306pub(super) const MISTRAL_EMBEDDING_WIDTHS: &[ModelWidth] = &[
307    ModelWidth {
308        model: crate::providers::mistral::embedding::MISTRAL_EMBED,
309        default: Some(1_024),
310        accepted: AcceptedWidths::Fixed,
311    },
312    ModelWidth {
313        model: "mistral-embed-2312",
314        default: Some(1_024),
315        accepted: AcceptedWidths::Fixed,
316    },
317    ModelWidth {
318        model: crate::providers::mistral::embedding::CODESTRAL_EMBED,
319        // Configurable with no documented native width, so a handle that
320        // names none reports 0 rather than inventing one.
321        default: None,
322        accepted: AcceptedWidths::Range {
323            // Mistral documents only a ceiling. The floor is rig's own
324            // "unknown" sentinel, which never reaches the wire.
325            min: 0,
326            max: 3_072,
327            requirement: "to be at most 3072 for Codestral Embed",
328        },
329    },
330    ModelWidth {
331        model: "codestral-embed-2505",
332        default: None,
333        accepted: AcceptedWidths::Range {
334            min: 0,
335            max: 3_072,
336            requirement: "to be at most 3072 for Codestral Embed",
337        },
338    },
339];
340
341/// Doubleword embedding widths: 32 through 4096, defaulting to 4096.
342/// Validate both bounds locally because out-of-range requests can be clamped or
343/// inconsistently rejected by the service.
344pub(super) const DOUBLEWORD_EMBEDDING_WIDTHS: &[ModelWidth] = &[ModelWidth {
345    model: crate::providers::doubleword::QWEN3_EMBEDDING_8B,
346    default: Some(4_096),
347    accepted: AcceptedWidths::Range {
348        min: 32,
349        max: 4_096,
350        requirement: "to be between 32 and 4096",
351    },
352}];
353
354/// Encode an authenticated JSON POST with whole-response framing.
355fn json_post(
356    provider: &OpenAIConfig,
357    path: &str,
358    deployment: Option<&str>,
359    body: &serde_json::Value,
360) -> Result<Encoded, EncodeError> {
361    json_post_to(provider, provider.uri(path, deployment), body)
362}
363
364/// [`json_post`] against an already-resolved URL, for the endpoints whose
365/// URL is derived rather than a fixed path under the base.
366fn json_post_to(
367    provider: &OpenAIConfig,
368    uri: String,
369    body: &serde_json::Value,
370) -> Result<Encoded, EncodeError> {
371    let bytes = serde_json::to_vec(body)?;
372    let builder = http::Request::post(uri).header("Content-Type", "application/json");
373    encoded(provider, builder, Body::Bytes(bytes))
374}
375
376/// The `GET` whose status is the answer, for the two endpoints that send no
377/// body: the model catalogue and the credential check.
378fn get(provider: &OpenAIConfig, path: &str) -> Result<Encoded, EncodeError> {
379    encoded(
380        provider,
381        http::Request::get(provider.uri(path, None)),
382        Body::empty(),
383    )
384}
385
386/// Authenticate and build a request, then apply its modality envelope hook.
387/// Return construction or hook errors. Use whole-response framing and the
388/// dialect's request-ID header.
389fn encoded(
390    provider: &OpenAIConfig,
391    builder: http::request::Builder,
392    body: Body,
393) -> Result<Encoded, EncodeError> {
394    let mut request = provider.authenticate(builder).body(body)?;
395    if let Some(envelope) = provider
396        .dialect
397        .quirks
398        .hooks
399        .and_then(|hooks| hooks.modality_envelope)
400    {
401        envelope(provider, &mut request)?;
402    }
403    Ok(Encoded::new(request, Framing::Whole)
404        .with_request_id_header(provider.dialect.request_id_header))
405}
406
407/// Refuse an embeddings request parameter the dialect does not accept.
408fn unsupported_parameter(provider: &str, parameter: &str) -> EncodeError {
409    EncodeError::request(format!(
410        "{provider} embeddings do not support the `{parameter}` parameter"
411    ))
412}
413
414/// The embeddings wire.
415#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
416pub struct Embeddings {
417    /// Which provider, and how to reach it.
418    pub provider: OpenAIConfig,
419    /// The embedding model.
420    pub model: String,
421    /// The width the caller asked for, when they named one rather than
422    /// taking the model's default.
423    pub ndims: Option<usize>,
424    /// The encoding the caller asked the provider to answer in.
425    pub encoding_format: Option<EncodingFormat>,
426    /// The end-user identifier the provider attributes the call to.
427    pub user: Option<String>,
428}
429
430impl Embeddings {
431    /// The embeddings wire for `model`.
432    pub fn new(provider: OpenAIConfig, model: impl Into<String>, ndims: Option<usize>) -> Self {
433        Self {
434            provider,
435            model: model.into(),
436            ndims,
437            encoding_format: None,
438            user: None,
439        }
440    }
441
442    /// Ask the provider to answer in `encoding_format`.
443    pub fn with_encoding_format(mut self, encoding_format: EncodingFormat) -> Self {
444        self.encoding_format = Some(encoding_format);
445        self
446    }
447
448    /// Attribute the call to an end user.
449    pub fn with_user(mut self, user: impl Into<String>) -> Self {
450        self.user = Some(user.into());
451        self
452    }
453
454    /// This model's width contract on this dialect, or `None` for a model
455    /// the dialect does not document.
456    fn model_width(&self) -> Option<&'static ModelWidth> {
457        self.provider
458            .dialect
459            .quirks
460            .embedding
461            .widths
462            .iter()
463            .find(|width| width.model == self.model)
464    }
465
466    /// Resolve width from the caller, dialect table, then shared model table.
467    /// Return zero when all are absent.
468    fn resolved_ndims(&self) -> usize {
469        self.ndims
470            .or_else(|| self.model_width().and_then(|width| width.default))
471            .or_else(|| model_dimensions_from_identifier(&self.model))
472            .unwrap_or_default()
473    }
474
475    /// Validate declared widths against the dialect's zero-width and model policies.
476    /// Return a request error for unsupported widths; unknown models are unchecked.
477    fn refuse_unhonourable_width(&self) -> Result<(), EncodeError> {
478        let quirks = &self.provider.dialect.quirks.embedding;
479        let provider = self.provider.dialect.name;
480        let invalid = |requirement, parameter| {
481            EncodeError::request(format!(
482                "{provider} embeddings require `{parameter}` {requirement}"
483            ))
484        };
485        // A dialect that reads no width field has nothing to refuse: the
486        // caller's number never reaches the wire, and the shared driver
487        // catches the disagreement against the reply instead.
488        let Some(parameter) = quirks.dimensions.name() else {
489            return Ok(());
490        };
491        let Some(declared) = self.ndims else {
492            return Ok(());
493        };
494        if declared == 0 {
495            return match quirks.refuse_zero_width {
496                Some(requirement) => Err(invalid(requirement, parameter)),
497                None => Ok(()),
498            };
499        }
500        // A model the dialect does not document: the caller's width is the
501        // only width there is, so it goes out unvalidated and the API rules.
502        let Some(width) = self.model_width() else {
503            return Ok(());
504        };
505        // Native widths require no truncation parameter.
506        if width.default == Some(declared) {
507            return Ok(());
508        }
509        match width.accepted {
510            AcceptedWidths::Fixed => Err(unsupported_parameter(provider, parameter)),
511            AcceptedWidths::Range { min, max, .. } if (min..=max).contains(&declared) => Ok(()),
512            AcceptedWidths::Range { requirement, .. } => Err(invalid(requirement, parameter)),
513        }
514    }
515
516    /// The width to send, in the field this dialect spells it with.
517    ///
518    /// OpenAI's legacy Ada model does not accept a width at all, and
519    /// `llama-server` reads no width field, so neither is sent one.
520    fn requested_width(&self) -> Option<(&'static str, usize)> {
521        let field = self.provider.dialect.quirks.embedding.dimensions.name()?;
522        if self.model == crate::providers::openai::embedding::TEXT_EMBEDDING_ADA_002 {
523            return None;
524        }
525        // Unknown widths are metadata sentinels, not request parameters.
526        let ndims = match self.resolved_ndims() {
527            0 => return None,
528            ndims => ndims,
529        };
530        // At a documented model's native width, send nothing: that width is
531        // what the model emits unasked, so the field would only restate the
532        // default and the vector is identical either way.
533        if self
534            .model_width()
535            .is_some_and(|width| width.default == Some(ndims))
536        {
537            return None;
538        }
539        Some((field, ndims))
540    }
541}
542
543/// Decode embedding vectors and usage, enforcing the dialect's usage requirement.
544#[derive(Default)]
545pub struct EmbeddingsDecoder {
546    /// Whether this dialect's reply must carry usage.
547    requires_usage: bool,
548    provider: &'static str,
549    model: String,
550}
551
552impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
553    type Event = CompatibleEmbeddingResponse;
554
555    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
556        classify_untyped_line(frame.as_str().as_bytes())
557    }
558
559    fn decode(
560        &mut self,
561        event: Self::Event,
562        out: Out<'id, Embedding>,
563    ) -> Result<Flow, ProviderError> {
564        if event.usage.is_none() && self.requires_usage {
565            return Err(ProviderError::Response(format!(
566                "{} embedding response omitted required usage",
567                self.provider
568            )));
569        }
570        let usage = event
571            .usage
572            .as_ref()
573            .map(Usage::to_normalized)
574            .unwrap_or_default();
575        let vectors = event.data.into_iter().map(|datum| {
576            datum
577                .embedding
578                .into_iter()
579                .filter_map(|number| number.as_f64())
580                .collect()
581        });
582        let model = if event.model.is_empty() {
583            self.model.clone()
584        } else {
585            event.model
586        };
587        Ok(out.end(embeddings::EmbeddingResponse {
588            model: Some(model),
589            usage,
590            ..embeddings::EmbeddingResponse::from_vectors(vectors)
591        }))
592    }
593}
594
595impl Wire for Embeddings {
596    type Op = Embedding;
597    type Payload = crate::wire::Encoded;
598    type Frame = crate::wire::WireFrame;
599    type Decoder<'id> = EmbeddingsDecoder;
600    type Reassembler = crate::wire::document::Unreassembled;
601
602    /// Only an explicit nonzero declaration permits reply-width mismatch checks.
603    fn describe(&self) -> Descriptor<'_> {
604        Descriptor::new(self.provider.dialect.name)
605            .model(self.model.as_str())
606            .capabilities(
607                Capabilities::embedding(
608                    self.provider.dialect.quirks.embedding.max_documents,
609                    self.resolved_ndims(),
610                )
611                .declaring(self.ndims),
612            )
613    }
614
615    fn encode(&self, request: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
616        let quirks = &self.provider.dialect.quirks.embedding;
617        // Base64 vectors are not decoded anywhere, so asking for them would
618        // answer 200 with a payload rig cannot read.
619        if self.encoding_format == Some(EncodingFormat::Base64) {
620            return Err(EncodeError::request(format!(
621                "Rig cannot decode {} embedding responses encoded as `base64`",
622                self.provider.dialect.name
623            )));
624        }
625        if self.encoding_format.is_some() && !quirks.supports_encoding_format {
626            return Err(unsupported_parameter(
627                self.provider.dialect.name,
628                "encoding_format",
629            ));
630        }
631        if self.user.is_some() && !quirks.supports_user {
632            return Err(unsupported_parameter(self.provider.dialect.name, "user"));
633        }
634        self.refuse_unhonourable_width()?;
635
636        let mut body = serde_json::json!({ "input": request });
637        let Some(object) = body.as_object_mut() else {
638            return Err(EncodeError::request(
639                "embedding request body must be an object",
640            ));
641        };
642        if quirks.sends_model_field {
643            object.insert("model".to_owned(), serde_json::json!(self.model));
644        }
645        if let Some((field, ndims)) = self.requested_width() {
646            object.insert(field.to_owned(), serde_json::json!(ndims));
647        }
648        if let Some(encoding_format) = self.encoding_format {
649            object.insert(
650                "encoding_format".to_owned(),
651                serde_json::to_value(encoding_format)?,
652            );
653        }
654        if let Some(user) = &self.user {
655            object.insert("user".to_owned(), serde_json::json!(user));
656        }
657
658        json_post(
659            &self.provider,
660            self.provider.dialect.quirks.embeddings_path,
661            self.provider.deployment(&self.model),
662            &body,
663        )
664    }
665
666    fn decoder<'id>(&self) -> Self::Decoder<'id> {
667        EmbeddingsDecoder {
668            requires_usage: self.provider.dialect.quirks.embedding.requires_usage,
669            provider: self.provider.dialect.name,
670            model: self.model.clone(),
671        }
672    }
673}
674
675/// Transcription requests encoded as multipart or input-audio JSON.
676#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
677pub struct Transcriptions {
678    /// Which provider, and how to reach it.
679    pub provider: OpenAIConfig,
680    /// Transcription model or Azure deployment identifier.
681    pub model: String,
682}
683
684impl Transcriptions {
685    /// The transcription wire for `model`.
686    pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
687        Self {
688            provider,
689            model: model.into(),
690        }
691    }
692
693    /// OpenAI's multipart upload: the audio as a file part beside the
694    /// per-request options.
695    fn multipart_body(&self, request: TranscriptionRequest) -> Result<Body, EncodeError> {
696        use crate::http_client::MultipartForm;
697        use crate::http_client::multipart::Part;
698
699        let mut form = MultipartForm::new();
700        // Azure addresses its deployment through the URL rather than the form.
701        if self.provider.deployment(&self.model).is_none() {
702            form = form.text("model", self.model.clone());
703        }
704        form = form.part(Part::bytes("file", request.data).filename(request.filename));
705        if let Some(language) = request.language {
706            form = form.text("language", language);
707        }
708        if let Some(prompt) = request.prompt {
709            form = form.text("prompt", prompt);
710        }
711        if let Some(temperature) = request.temperature {
712            form = form.text("temperature", temperature.to_string());
713        }
714        if let Some(additional_params) = request.additional_params {
715            for (name, value) in additional_params_object(&additional_params)? {
716                // Form strings must not acquire JSON quotation marks.
717                let value = match value {
718                    serde_json::Value::String(value) => value.clone(),
719                    other => other.to_string(),
720                };
721                form = form.text(name.clone(), value);
722            }
723        }
724        Ok(Body::Multipart(form))
725    }
726
727    /// Encode base64 audio under `input_audio`, inferring format from the filename.
728    /// Reject top-level prompts and non-object additional parameters.
729    fn input_audio_body(&self, request: TranscriptionRequest) -> Result<Body, EncodeError> {
730        use base64::Engine;
731
732        if request.prompt.is_some() {
733            return Err(EncodeError::request(std::io::Error::new(
734                std::io::ErrorKind::InvalidInput,
735                "OpenRouter STT does not support a top-level prompt field. \
736                     Provider-specific prompt options can be passed via `additional_params`. \
737                     Example: {\"provider\": {\"options\": {\"<provider>\": {\"prompt\": \"<text>\"}}}}",
738            )));
739        }
740
741        let mut body = serde_json::Map::new();
742        body.insert("model".to_owned(), serde_json::json!(self.model));
743        body.insert(
744            "input_audio".to_owned(),
745            serde_json::json!({
746                "data": base64::engine::general_purpose::STANDARD.encode(&request.data),
747                "format": audio_format_of(&request.filename),
748            }),
749        );
750        if let Some(language) = request.language {
751            body.insert("language".to_owned(), serde_json::json!(language));
752        }
753        if let Some(temperature) = request.temperature {
754            body.insert("temperature".to_owned(), serde_json::json!(temperature));
755        }
756        if let Some(additional_params) = request.additional_params {
757            for (name, value) in additional_params_object(&additional_params)? {
758                body.insert(name.clone(), value.clone());
759            }
760        }
761        Ok(Body::Bytes(serde_json::to_vec(
762            &serde_json::Value::Object(body),
763        )?))
764    }
765}
766
767/// A transcription request's `additional_params`, as an object.
768fn additional_params_object(
769    params: &serde_json::Value,
770) -> Result<&serde_json::Map<String, serde_json::Value>, EncodeError> {
771    params.as_object().ok_or_else(|| {
772        EncodeError::request(std::io::Error::new(
773            std::io::ErrorKind::InvalidInput,
774            "additional transcription parameters must be a JSON object",
775        ))
776    })
777}
778
779/// Infer the audio container from a case-insensitive filename extension.
780/// Unknown or absent extensions default to `wav`.
781fn audio_format_of(filename: &str) -> &'static str {
782    let extension = std::path::Path::new(filename)
783        .extension()
784        .and_then(std::ffi::OsStr::to_str)
785        .map(str::to_ascii_lowercase);
786    match extension.as_deref() {
787        Some("mp3") => "mp3",
788        Some("flac") => "flac",
789        Some("m4a") => "m4a",
790        Some("ogg") => "ogg",
791        Some("webm") => "webm",
792        Some("aac") => "aac",
793        _ => "wav",
794    }
795}
796
797/// The transcription decoder.
798#[derive(Default)]
799pub struct TranscriptionsDecoder;
800
801impl<'id> Decoder<'id, Transcription> for TranscriptionsDecoder {
802    type Event = crate::providers::openai::transcription::TranscriptionResponse;
803
804    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
805        classify_untyped_line(frame.as_str().as_bytes())
806    }
807
808    fn decode(
809        &mut self,
810        event: Self::Event,
811        out: Out<'id, Transcription>,
812    ) -> Result<Flow, ProviderError> {
813        Ok(out.end(event.normalize()?))
814    }
815}
816
817impl Wire for Transcriptions {
818    type Op = Transcription;
819    type Payload = crate::wire::Encoded;
820    type Frame = crate::wire::WireFrame;
821    type Decoder<'id> = TranscriptionsDecoder;
822    type Reassembler = crate::wire::document::Unreassembled;
823
824    fn describe(&self) -> Descriptor<'_> {
825        Descriptor::new(self.provider.dialect.name).model(self.model.as_str())
826    }
827
828    fn encode(&self, request: TranscriptionRequest, _mode: Mode) -> Result<Encoded, EncodeError> {
829        let uri = self
830            .provider
831            .modality_uri(
832                "transcription",
833                self.provider.dialect.quirks.transcription_path,
834                &self.model,
835            )
836            .map_err(EncodeError::request)?;
837        let builder = http::Request::post(uri);
838        let (builder, body) = match self.provider.dialect.quirks.transcription_body {
839            TranscriptionBody::Multipart => (builder, self.multipart_body(request)?),
840            TranscriptionBody::InputAudioJson => (
841                builder.header(http::header::CONTENT_TYPE, "application/json"),
842                self.input_audio_body(request)?,
843            ),
844        };
845        encoded(&self.provider, builder, body)
846    }
847
848    fn decoder<'id>(&self) -> Self::Decoder<'id> {
849        TranscriptionsDecoder
850    }
851}
852
853/// The image-generation wire.
854#[cfg(feature = "image")]
855#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
856pub struct Images {
857    /// Which provider, and how to reach it.
858    pub provider: OpenAIConfig,
859    /// The image model.
860    pub model: String,
861}
862
863#[cfg(feature = "image")]
864impl Images {
865    /// The image-generation wire for `model`.
866    pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
867        Self {
868            provider,
869            model: model.into(),
870        }
871    }
872}
873
874/// The image-generation decoder.
875///
876/// Carries the dialect's [`ImageBody`] because the request shape names the
877/// reply shape: most of this family answers with a JSON envelope, and
878/// Hugging Face's router answers with the image bytes themselves.
879#[cfg(feature = "image")]
880#[derive(Default)]
881pub struct ImagesDecoder {
882    /// Which reply shape this dialect answers with.
883    body: ImageBody,
884}
885
886/// One generated image, as every dialect on this wire returns it.
887#[cfg(feature = "image")]
888#[derive(Debug, Clone, Serialize, Deserialize)]
889pub struct ImageDatum {
890    /// The image, base64-encoded.
891    pub b64_json: String,
892}
893
894/// Base64 image represented as a string or an object with an `image` field.
895#[cfg(feature = "image")]
896#[derive(Debug, Clone, Serialize, Deserialize)]
897#[serde(untagged)]
898pub enum ImagesReplyImage {
899    /// Hyperbolic: `{"image": "<base64>"}`.
900    Keyed {
901        /// The image, base64-encoded.
902        image: String,
903    },
904    /// Venice: the base64 payload itself.
905    Bare(String),
906}
907
908#[cfg(feature = "image")]
909impl ImagesReplyImage {
910    /// The image's base64 payload, whichever form the dialect sent.
911    pub fn base64(&self) -> &str {
912        match self {
913            Self::Keyed { image } => image,
914            Self::Bare(image) => image,
915        }
916    }
917}
918
919/// JSON image reply accepting `data` and `images` arrays with optional metadata.
920#[cfg(feature = "image")]
921#[derive(Debug, Clone, Serialize, Deserialize)]
922pub struct ImagesReply {
923    /// The generated images, as OpenAI and xAI key them.
924    #[serde(default)]
925    pub data: Vec<ImageDatum>,
926    /// The generated images, as Hyperbolic and Venice key them.
927    #[serde(default)]
928    pub images: Vec<ImagesReplyImage>,
929    /// Whatever else the dialect sent (`created`, Venice's `id`/`timing`, and
930    /// any field this build does not model), so the raw payload loses
931    /// nothing.
932    #[serde(flatten)]
933    pub extra: serde_json::Map<String, serde_json::Value>,
934}
935
936#[cfg(feature = "image")]
937impl ImagesReply {
938    /// The first image's base64 payload, whichever key the dialect used.
939    pub fn first_base64(&self) -> Option<&str> {
940        self.data
941            .first()
942            .map(|image| image.b64_json.as_str())
943            .or_else(|| self.images.first().map(ImagesReplyImage::base64))
944            .filter(|encoded| !encoded.is_empty())
945    }
946}
947
948/// Generated images in a decoded JSON envelope or as raw bytes.
949#[cfg(feature = "image")]
950#[derive(Debug, Clone)]
951pub enum ImagesEvent {
952    /// A JSON envelope, as OpenAI, xAI and Hyperbolic answer.
953    Json(ImagesReply),
954    /// The image bytes themselves, with no envelope at all.
955    Raw(Vec<u8>),
956}
957
958#[cfg(feature = "image")]
959impl<'id> Decoder<'id, crate::operation::ImageGeneration> for ImagesDecoder {
960    type Event = ImagesEvent;
961
962    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
963        match self.body {
964            // Preserve raw image bytes even when framing recognized valid UTF-8.
965            ImageBody::HuggingFace => WireEvent::Known(ImagesEvent::Raw(match frame {
966                WireFrame::Text(text) => text.into_bytes(),
967                WireFrame::Bytes(bytes) => bytes,
968            })),
969            ImageBody::OpenAi | ImageBody::Xai | ImageBody::Hyperbolic | ImageBody::Venice => {
970                classify_untyped_line(frame.as_str().as_bytes()).map(ImagesEvent::Json)
971            }
972        }
973    }
974
975    fn decode(
976        &mut self,
977        event: Self::Event,
978        out: Out<'id, crate::operation::ImageGeneration>,
979    ) -> Result<Flow, ProviderError> {
980        use crate::image_generation::ImageGenerationResponse;
981        use base64::Engine;
982
983        let reply = match event {
984            // The image is already the payload, and the reply is not a
985            // document, so `raw` stays null rather than restating the bytes.
986            ImagesEvent::Raw(image) => {
987                return Ok(out.end(ImageGenerationResponse::new(image)));
988            }
989            ImagesEvent::Json(reply) => reply,
990        };
991        let Some(encoded) = reply.first_base64() else {
992            return Err(ProviderError::Response("missing image data".to_owned()));
993        };
994        let image = match base64::prelude::BASE64_STANDARD.decode(encoded) {
995            Ok(image) => image,
996            Err(error) => {
997                return Err(ProviderError::Response(error.to_string()));
998            }
999        };
1000        Ok(out.end(ImageGenerationResponse::new(image)))
1001    }
1002}
1003
1004#[cfg(feature = "image")]
1005impl Wire for Images {
1006    type Op = crate::operation::ImageGeneration;
1007    type Payload = crate::wire::Encoded;
1008    type Frame = crate::wire::WireFrame;
1009    type Decoder<'id> = ImagesDecoder;
1010    type Reassembler = crate::wire::document::Unreassembled;
1011
1012    fn describe(&self) -> Descriptor<'_> {
1013        Descriptor::new(self.provider.dialect.name).model(self.model.as_str())
1014    }
1015
1016    fn encode(
1017        &self,
1018        request: crate::image_generation::ImageGenerationRequest,
1019        _mode: Mode,
1020    ) -> Result<Encoded, EncodeError> {
1021        let mut body = match self.provider.dialect.quirks.image_body {
1022            // `response_format` is deliberately absent: it is no longer part
1023            // of OpenAI's request schema, which rejects it before it even
1024            // looks at the model. A compatible endpoint that still takes the
1025            // field gets it through `additional_params`.
1026            ImageBody::OpenAi => serde_json::json!({
1027                "model": self.model,
1028                "prompt": request.prompt,
1029                "size": format!("{}x{}", request.width, request.height),
1030            }),
1031            // xAI takes no `size` and answers with a URL unless asked for
1032            // base64, which is the only form this wire decodes.
1033            ImageBody::Xai => serde_json::json!({
1034                "model": self.model,
1035                "prompt": request.prompt,
1036                "response_format": "b64_json",
1037                "aspect_ratio": "1:1",
1038            }),
1039            ImageBody::Hyperbolic => serde_json::json!({
1040                "model_name": self.model,
1041                "prompt": request.prompt,
1042                "height": request.height,
1043                "width": request.width,
1044            }),
1045            ImageBody::Venice => serde_json::json!({
1046                "model": self.model,
1047                "prompt": request.prompt,
1048                "width": request.width,
1049                "height": request.height,
1050            }),
1051            // The model is addressed through the URL, not the body.
1052            ImageBody::HuggingFace => serde_json::json!({
1053                "inputs": request.prompt,
1054                "parameters": {
1055                    "width": request.width,
1056                    "height": request.height,
1057                },
1058            }),
1059        };
1060        // Merged last, so a caller can reach the endpoint's other parameters
1061        // and override what is derived above.
1062        if let Some(additional_params) = request.additional_params {
1063            crate::json_utils::merge_inplace(&mut body, additional_params);
1064        }
1065
1066        let uri = self
1067            .provider
1068            .modality_uri(
1069                "image generation",
1070                self.provider.dialect.quirks.image_generation_path,
1071                &self.model,
1072            )
1073            .map_err(EncodeError::request)?;
1074        json_post_to(&self.provider, uri, &body)
1075    }
1076
1077    fn decoder<'id>(&self) -> Self::Decoder<'id> {
1078        ImagesDecoder {
1079            body: self.provider.dialect.quirks.image_body,
1080        }
1081    }
1082}
1083
1084/// The speech wire.
1085#[cfg(feature = "audio")]
1086#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1087pub struct Speech {
1088    /// Which provider, and how to reach it.
1089    pub provider: OpenAIConfig,
1090    /// The speech model.
1091    pub model: String,
1092}
1093
1094#[cfg(feature = "audio")]
1095impl Speech {
1096    /// The speech wire for `model`.
1097    pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
1098        Self {
1099            provider,
1100            model: model.into(),
1101        }
1102    }
1103}
1104
1105/// Decode raw audio or the dialect's base64 JSON envelope.
1106#[cfg(feature = "audio")]
1107#[derive(Default)]
1108pub struct SpeechDecoder {
1109    /// Which reply shape this dialect answers with.
1110    body: SpeechBody,
1111}
1112
1113/// Hyperbolic's speech reply: base64 in a JSON envelope rather than the
1114/// audio bytes themselves.
1115#[cfg(feature = "audio")]
1116#[derive(Debug, Clone, Serialize, Deserialize)]
1117pub struct SpeechReply {
1118    /// The audio, base64-encoded.
1119    pub audio: String,
1120}
1121
1122#[cfg(feature = "audio")]
1123impl<'id> Decoder<'id, crate::operation::AudioGeneration> for SpeechDecoder {
1124    type Event = Vec<u8>;
1125
1126    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
1127        // Interpretation selects raw audio or JSON decoding from the dialect.
1128        WireEvent::Known(match frame {
1129            WireFrame::Text(text) => text.into_bytes(),
1130            WireFrame::Bytes(bytes) => bytes,
1131        })
1132    }
1133
1134    fn decode(
1135        &mut self,
1136        event: Self::Event,
1137        out: Out<'id, crate::operation::AudioGeneration>,
1138    ) -> Result<Flow, ProviderError> {
1139        use base64::Engine;
1140
1141        let audio = match self.body {
1142            // OpenAI and xAI answer with the audio itself.
1143            SpeechBody::OpenAi | SpeechBody::Xai => event,
1144            // Hyperbolic wraps it, base64-encoded, in a JSON envelope.
1145            SpeechBody::Hyperbolic => {
1146                let reply = match serde_json::from_slice::<SpeechReply>(&event) {
1147                    Ok(reply) => reply,
1148                    Err(error) => {
1149                        return Err(ProviderError::Response(error.to_string()));
1150                    }
1151                };
1152                match base64::prelude::BASE64_STANDARD.decode(&reply.audio) {
1153                    Ok(audio) => audio,
1154                    Err(error) => {
1155                        return Err(ProviderError::Response(error.to_string()));
1156                    }
1157                }
1158            }
1159        };
1160        Ok(out.end(crate::audio_generation::AudioGenerationResponse::new(audio)))
1161    }
1162}
1163
1164#[cfg(feature = "audio")]
1165impl Wire for Speech {
1166    type Op = crate::operation::AudioGeneration;
1167    type Payload = crate::wire::Encoded;
1168    type Frame = crate::wire::WireFrame;
1169    type Decoder<'id> = SpeechDecoder;
1170    type Reassembler = crate::wire::document::Unreassembled;
1171
1172    fn describe(&self) -> Descriptor<'_> {
1173        Descriptor::new(self.provider.dialect.name).model(self.model.as_str())
1174    }
1175
1176    fn encode(
1177        &self,
1178        request: crate::audio_generation::AudioGenerationRequest,
1179        _mode: Mode,
1180    ) -> Result<Encoded, EncodeError> {
1181        let mut body = match self.provider.dialect.quirks.speech_body {
1182            SpeechBody::OpenAi => serde_json::json!({
1183                "model": self.model,
1184                "input": request.text,
1185                "voice": request.voice,
1186                "speed": request.speed,
1187            }),
1188            // xAI requires a default voice when the caller leaves it empty.
1189            SpeechBody::Xai => serde_json::json!({
1190                "text": request.text,
1191                "voice_id": if request.voice.is_empty() { "eve" } else { request.voice.as_str() },
1192                "language": "en",
1193            }),
1194            // Hyperbolic addresses this endpoint by language, so the
1195            // identifier the caller passes as the model IS the language tag.
1196            SpeechBody::Hyperbolic => serde_json::json!({
1197                "language": self.model,
1198                "speaker": request.voice,
1199                "text": request.text,
1200                "speed": request.speed,
1201            }),
1202        };
1203        // Caller parameters take precedence, including response format and instructions.
1204        if let Some(additional_params) = request.additional_params {
1205            crate::json_utils::merge_inplace(&mut body, additional_params);
1206        }
1207
1208        // Azure versions its speech endpoint separately from every other
1209        // route, so this one request carries its own `api-version`.
1210        let uri = self.provider.uri_versioned(
1211            self.provider.dialect.quirks.audio_generation_path,
1212            self.provider.deployment(&self.model),
1213            self.provider.speech_api_version(),
1214        );
1215        json_post_to(&self.provider, uri, &body)
1216    }
1217
1218    fn decoder<'id>(&self) -> Self::Decoder<'id> {
1219        SpeechDecoder {
1220            body: self.provider.dialect.quirks.speech_body,
1221        }
1222    }
1223}
1224
1225/// The model-listing wire.
1226#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1227pub struct Models {
1228    /// Which provider, and how to reach it.
1229    pub provider: OpenAIConfig,
1230}
1231
1232impl Models {
1233    /// The model-listing wire.
1234    pub fn new(provider: OpenAIConfig) -> Self {
1235        Self { provider }
1236    }
1237}
1238
1239/// Model-list entry with required `id` and optional metadata.
1240/// Context length prefers `context_window`, then `context_length`, then
1241/// `max_context_length`. Top-level output limits take precedence over `top_provider`.
1242#[derive(Debug, Deserialize)]
1243pub struct ModelEntry {
1244    pub id: String,
1245    #[serde(default)]
1246    pub name: Option<String>,
1247    #[serde(default)]
1248    pub description: Option<String>,
1249    /// Mistral labels the model kind `type` (`base`, `fine-tuned`).
1250    #[serde(default, rename = "type")]
1251    pub kind: Option<String>,
1252    #[serde(default)]
1253    pub created: Option<u64>,
1254    #[serde(default)]
1255    pub owned_by: Option<String>,
1256    #[serde(default)]
1257    pub context_window: Option<u32>,
1258    #[serde(default)]
1259    pub context_length: Option<u32>,
1260    #[serde(default)]
1261    pub max_context_length: Option<u32>,
1262    #[serde(default)]
1263    pub max_completion_tokens: Option<u32>,
1264    #[serde(default)]
1265    pub top_provider: Option<TopProvider>,
1266}
1267
1268/// OpenRouter's per-entry routing block. Only the output ceiling is read;
1269/// the rest of the block is routing detail [`ModelInfo`] has no slot for.
1270#[derive(Debug, Deserialize)]
1271pub struct TopProvider {
1272    #[serde(default)]
1273    pub max_completion_tokens: Option<u32>,
1274}
1275
1276impl From<ModelEntry> for ModelInfo {
1277    fn from(entry: ModelEntry) -> Self {
1278        let mut model = ModelInfo::from_id(entry.id);
1279        model.name = entry.name;
1280        model.description = entry.description;
1281        model.r#type = entry.kind;
1282        model.created_at = entry.created;
1283        model.owned_by = entry.owned_by;
1284        model.context_length = entry
1285            .context_window
1286            .or(entry.context_length)
1287            .or(entry.max_context_length);
1288        model.max_output_tokens = entry.max_completion_tokens.or_else(|| {
1289            entry
1290                .top_provider
1291                .and_then(|provider| provider.max_completion_tokens)
1292        });
1293        model
1294    }
1295}
1296
1297/// The `{ "data": [...] }` envelope.
1298#[derive(Debug, Deserialize)]
1299pub struct ModelsReply {
1300    #[serde(default)]
1301    pub data: Vec<ModelEntry>,
1302}
1303
1304/// Decode a complete model catalogue without pagination.
1305#[derive(Default)]
1306pub struct ModelsDecoder;
1307
1308impl<'id> Decoder<'id, ModelListing> for ModelsDecoder {
1309    type Event = ModelsReply;
1310
1311    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
1312        classify_untyped_line(frame.as_str().as_bytes())
1313    }
1314
1315    fn decode(
1316        &mut self,
1317        event: Self::Event,
1318        out: Out<'id, ModelListing>,
1319    ) -> Result<Flow, ProviderError> {
1320        let models = event.data.into_iter().map(ModelInfo::from).collect();
1321        Ok(out.end(ModelPage {
1322            models: ModelList::new(models),
1323            next: None,
1324        }))
1325    }
1326}
1327
1328impl Wire for Models {
1329    type Op = ModelListing;
1330    type Payload = crate::wire::Encoded;
1331    type Frame = crate::wire::WireFrame;
1332    type Decoder<'id> = ModelsDecoder;
1333    type Reassembler = crate::wire::document::Unreassembled;
1334
1335    fn describe(&self) -> Descriptor<'_> {
1336        Descriptor::new(self.provider.dialect.name)
1337    }
1338
1339    /// The catalogue is one page, so a cursor never follows.
1340    fn encode(&self, _cursor: Option<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
1341        get(&self.provider, self.provider.dialect.quirks.models_path)
1342    }
1343
1344    fn decoder<'id>(&self) -> Self::Decoder<'id> {
1345        ModelsDecoder
1346    }
1347}
1348
1349/// Rerank documents with `{model, query, documents, top_n}` requests.
1350/// The dialect must configure a nonempty reranking path.
1351#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1352pub struct Rerank {
1353    /// Which provider, and how to reach it.
1354    pub provider: OpenAIConfig,
1355    /// The reranker model.
1356    pub model: String,
1357    /// Return only the `top_n` highest-scoring documents, when the caller
1358    /// asked for a cut.
1359    pub top_n: Option<usize>,
1360}
1361
1362impl Rerank {
1363    /// The rerank wire for `model`.
1364    pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
1365        Self {
1366            provider,
1367            model: model.into(),
1368            top_n: None,
1369        }
1370    }
1371
1372    /// Ask the server to return only the `top_n` highest-scoring documents.
1373    pub fn with_top_n(mut self, top_n: usize) -> Self {
1374        self.top_n = Some(top_n);
1375        self
1376    }
1377}
1378
1379/// Scored input document accepting `relevance_score` or `score` as its score key.
1380#[derive(Debug, Clone, Serialize, Deserialize)]
1381pub struct RerankResultEntry {
1382    /// Which input document this scored.
1383    pub index: usize,
1384    /// The score.
1385    #[serde(alias = "score")]
1386    pub relevance_score: f64,
1387    /// Present only on servers that echo the document back; llama.cpp does
1388    /// not on this path.
1389    #[serde(default, alias = "text")]
1390    pub document: Option<String>,
1391}
1392
1393/// What a rerank reply reports besides its ranking.
1394#[derive(Debug, Clone, Serialize, Deserialize, Default)]
1395pub struct RerankUsage {
1396    /// Tokens the query and documents cost.
1397    #[serde(default)]
1398    pub prompt_tokens: u64,
1399    /// Total tokens, as the provider reported them.
1400    #[serde(default)]
1401    pub total_tokens: u64,
1402}
1403
1404/// The rerank reply.
1405#[derive(Debug, Clone, Serialize, Deserialize)]
1406pub struct RerankReply {
1407    /// The model the server ranked with, when it named one.
1408    #[serde(default)]
1409    pub model: Option<String>,
1410    /// The ranking.
1411    pub results: Vec<RerankResultEntry>,
1412    /// What it cost.
1413    #[serde(default)]
1414    pub usage: Option<RerankUsage>,
1415}
1416
1417/// The rerank decoder.
1418#[derive(Default)]
1419pub struct RerankDecoder;
1420
1421impl<'id> Decoder<'id, RerankOp> for RerankDecoder {
1422    type Event = RerankReply;
1423
1424    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
1425        classify_untyped_line(frame.as_str().as_bytes())
1426    }
1427
1428    fn decode(
1429        &mut self,
1430        event: Self::Event,
1431        out: Out<'id, RerankOp>,
1432    ) -> Result<Flow, ProviderError> {
1433        let usage = event
1434            .usage
1435            .map(|usage| crate::completion::Usage {
1436                input_tokens: Some(usage.prompt_tokens),
1437                total_tokens: Some(usage.total_tokens),
1438                ..Default::default()
1439            })
1440            .unwrap_or_default();
1441        let results = event
1442            .results
1443            .into_iter()
1444            .map(|result| crate::rerank::RerankResult {
1445                index: result.index,
1446                document: result.document,
1447                relevance_score: result.relevance_score,
1448            })
1449            .collect();
1450        Ok(out.end(crate::rerank::RerankResponse {
1451            model: event.model,
1452            usage,
1453            ..crate::rerank::RerankResponse::new(results)
1454        }))
1455    }
1456}
1457
1458impl Wire for Rerank {
1459    type Op = RerankOp;
1460    type Payload = crate::wire::Encoded;
1461    type Frame = crate::wire::WireFrame;
1462    type Decoder<'id> = RerankDecoder;
1463    type Reassembler = crate::wire::document::Unreassembled;
1464
1465    fn describe(&self) -> Descriptor<'_> {
1466        Descriptor::new(self.provider.dialect.name)
1467            .model(self.model.as_str())
1468            .capabilities(Capabilities::rerank(
1469                self.provider.dialect.quirks.rerank.max_documents,
1470            ))
1471    }
1472
1473    fn encode(
1474        &self,
1475        request: crate::operation::RerankRequest,
1476        _mode: Mode,
1477    ) -> Result<Encoded, EncodeError> {
1478        let quirks = &self.provider.dialect.quirks.rerank;
1479        // An empty path explicitly disables reranking.
1480        if quirks.path.is_empty() {
1481            return Err(EncodeError::request(format!(
1482                "{} offers no reranking endpoint",
1483                self.provider.dialect.name
1484            )));
1485        }
1486        let mut body = serde_json::json!({
1487            "query": request.query,
1488            "documents": request.documents,
1489        });
1490        let Some(object) = body.as_object_mut() else {
1491            return Err(EncodeError::request(
1492                "rerank request body must be an object",
1493            ));
1494        };
1495        if quirks.sends_model_field {
1496            object.insert("model".to_owned(), serde_json::json!(self.model));
1497        }
1498        if let Some(top_n) = self.top_n {
1499            object.insert("top_n".to_owned(), serde_json::json!(top_n));
1500        }
1501
1502        json_post(
1503            &self.provider,
1504            quirks.path,
1505            self.provider.deployment(&self.model),
1506            &body,
1507        )
1508    }
1509
1510    fn decoder<'id>(&self) -> Self::Decoder<'id> {
1511        RerankDecoder
1512    }
1513}
1514
1515/// The credential-check wire: a `GET` whose status is the answer.
1516#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1517pub struct Verify {
1518    /// Which provider, and how to reach it.
1519    pub provider: OpenAIConfig,
1520}
1521
1522impl Verify {
1523    /// The credential-check wire.
1524    pub fn new(provider: OpenAIConfig) -> Self {
1525        Self { provider }
1526    }
1527}
1528
1529pub use crate::operation::VerifyDecoder;
1530
1531impl Wire for Verify {
1532    type Op = VerifyOp;
1533    type Payload = crate::wire::Encoded;
1534    type Frame = crate::wire::WireFrame;
1535    type Decoder<'id> = VerifyDecoder;
1536    type Reassembler = crate::wire::document::Unreassembled;
1537
1538    fn describe(&self) -> Descriptor<'_> {
1539        Descriptor::new(self.provider.dialect.name)
1540    }
1541
1542    fn encode(&self, _request: (), _mode: Mode) -> Result<Encoded, EncodeError> {
1543        let path = self.provider.dialect.quirks.verify_path;
1544        if path.is_empty() {
1545            return Err(EncodeError::request(format!(
1546                "{} offers no endpoint that checks a credential without consuming tokens",
1547                self.provider.dialect.name
1548            )));
1549        }
1550        get(&self.provider, path)
1551    }
1552
1553    fn decoder<'id>(&self) -> Self::Decoder<'id> {
1554        VerifyDecoder
1555    }
1556}
1557
1558#[cfg(test)]
1559mod tests;