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::completion::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
30#[cfg(feature = "image")]
31use super::ImageBody;
32#[cfg(feature = "audio")]
33use super::SpeechBody;
34use super::{AcceptedWidths, ModelWidth, OpenAIConfig, TranscriptionBody};
35
36/// Encode an authenticated JSON POST with whole-response framing.
37fn json_post(
38    provider: &OpenAIConfig,
39    path: &str,
40    deployment: Option<&str>,
41    body: &serde_json::Value,
42) -> Result<Encoded, EncodeError> {
43    json_post_to(provider, provider.uri(path, deployment), body)
44}
45
46/// [`json_post`] against an already-resolved URL, for the endpoints whose
47/// URL is derived rather than a fixed path under the base.
48fn json_post_to(
49    provider: &OpenAIConfig,
50    uri: String,
51    body: &serde_json::Value,
52) -> Result<Encoded, EncodeError> {
53    let bytes = serde_json::to_vec(body)?;
54    let builder = http::Request::post(uri).header("Content-Type", "application/json");
55    encoded(provider, builder, Body::Bytes(bytes))
56}
57
58/// The `GET` whose status is the answer, for the two endpoints that send no
59/// body: the model catalogue and the credential check.
60fn get(provider: &OpenAIConfig, path: &str) -> Result<Encoded, EncodeError> {
61    encoded(
62        provider,
63        http::Request::get(provider.uri(path, None)),
64        Body::empty(),
65    )
66}
67
68/// Authenticate and build a request, then apply its modality envelope hook.
69/// Return construction or hook errors. Use whole-response framing and the
70/// dialect's request-ID header.
71fn encoded(
72    provider: &OpenAIConfig,
73    builder: http::request::Builder,
74    body: Body,
75) -> Result<Encoded, EncodeError> {
76    let mut request = provider.authenticate(builder).body(body)?;
77    if let Some(envelope) = provider
78        .dialect
79        .quirks
80        .hooks
81        .and_then(|hooks| hooks.modality_envelope)
82    {
83        envelope(provider, &mut request)?;
84    }
85    Ok(Encoded::new(request, Framing::Whole)
86        .with_request_id_header(provider.dialect.request_id_header))
87}
88
89/// Refuse an embeddings request parameter the dialect does not accept.
90fn unsupported_parameter(provider: &str, parameter: &str) -> EncodeError {
91    EncodeError::request(format!(
92        "{provider} embeddings do not support the `{parameter}` parameter"
93    ))
94}
95
96/// The embeddings wire.
97#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
98pub struct Embeddings {
99    /// Which provider, and how to reach it.
100    pub provider: OpenAIConfig,
101    /// The embedding model.
102    pub model: String,
103    /// The width the caller asked for, when they named one rather than
104    /// taking the model's default.
105    pub ndims: Option<usize>,
106    /// The encoding the caller asked the provider to answer in.
107    pub encoding_format: Option<EncodingFormat>,
108    /// The end-user identifier the provider attributes the call to.
109    pub user: Option<String>,
110}
111
112impl Embeddings {
113    /// The embeddings wire for `model`.
114    pub fn new(provider: OpenAIConfig, model: impl Into<String>, ndims: Option<usize>) -> Self {
115        Self {
116            provider,
117            model: model.into(),
118            ndims,
119            encoding_format: None,
120            user: None,
121        }
122    }
123
124    /// Ask the provider to answer in `encoding_format`.
125    pub fn with_encoding_format(mut self, encoding_format: EncodingFormat) -> Self {
126        self.encoding_format = Some(encoding_format);
127        self
128    }
129
130    /// Attribute the call to an end user.
131    pub fn with_user(mut self, user: impl Into<String>) -> Self {
132        self.user = Some(user.into());
133        self
134    }
135
136    /// This model's width contract on this dialect, or `None` for a model
137    /// the dialect does not document.
138    fn model_width(&self) -> Option<&'static ModelWidth> {
139        self.provider
140            .dialect
141            .quirks
142            .embedding
143            .widths
144            .iter()
145            .find(|width| width.model == self.model)
146    }
147
148    /// Resolve width from the caller, dialect table, then shared model table.
149    /// Return zero when all are absent.
150    fn resolved_ndims(&self) -> usize {
151        self.ndims
152            .or_else(|| self.model_width().and_then(|width| width.default))
153            .or_else(|| model_dimensions_from_identifier(&self.model))
154            .unwrap_or_default()
155    }
156
157    /// Validate declared widths against the dialect's zero-width and model policies.
158    /// Return a request error for unsupported widths; unknown models are unchecked.
159    fn refuse_unhonourable_width(&self) -> Result<(), EncodeError> {
160        let quirks = &self.provider.dialect.quirks.embedding;
161        let provider = self.provider.dialect.name;
162        let invalid = |requirement, parameter| {
163            EncodeError::request(format!(
164                "{provider} embeddings require `{parameter}` {requirement}"
165            ))
166        };
167        // A dialect that reads no width field has nothing to refuse: the
168        // caller's number never reaches the wire, and the shared driver
169        // catches the disagreement against the reply instead.
170        let Some(parameter) = quirks.dimensions.name() else {
171            return Ok(());
172        };
173        let Some(declared) = self.ndims else {
174            return Ok(());
175        };
176        if declared == 0 {
177            return match quirks.refuse_zero_width {
178                Some(requirement) => Err(invalid(requirement, parameter)),
179                None => Ok(()),
180            };
181        }
182        // A model the dialect does not document: the caller's width is the
183        // only width there is, so it goes out unvalidated and the API rules.
184        let Some(width) = self.model_width() else {
185            return Ok(());
186        };
187        // Native widths require no truncation parameter.
188        if width.default == Some(declared) {
189            return Ok(());
190        }
191        match width.accepted {
192            AcceptedWidths::Fixed => Err(unsupported_parameter(provider, parameter)),
193            AcceptedWidths::Range { min, max, .. } if (min..=max).contains(&declared) => Ok(()),
194            AcceptedWidths::Range { requirement, .. } => Err(invalid(requirement, parameter)),
195        }
196    }
197
198    /// The width to send, in the field this dialect spells it with.
199    ///
200    /// OpenAI's legacy Ada model does not accept a width at all, and
201    /// `llama-server` reads no width field, so neither is sent one.
202    fn requested_width(&self) -> Option<(&'static str, usize)> {
203        let field = self.provider.dialect.quirks.embedding.dimensions.name()?;
204        if self.model == crate::providers::openai::embedding::TEXT_EMBEDDING_ADA_002 {
205            return None;
206        }
207        // Unknown widths are metadata sentinels, not request parameters.
208        let ndims = match self.resolved_ndims() {
209            0 => return None,
210            ndims => ndims,
211        };
212        // At a documented model's native width, send nothing: that width is
213        // what the model emits unasked, so the field would only restate the
214        // default and the vector is identical either way.
215        if self
216            .model_width()
217            .is_some_and(|width| width.default == Some(ndims))
218        {
219            return None;
220        }
221        Some((field, ndims))
222    }
223}
224
225/// Decode embedding vectors and usage, enforcing the dialect's usage requirement.
226#[derive(Default)]
227pub struct EmbeddingsDecoder {
228    /// Whether this dialect's reply must carry usage.
229    requires_usage: bool,
230    provider: &'static str,
231    model: String,
232}
233
234impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
235    type Event = CompatibleEmbeddingResponse;
236
237    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
238        classify_untyped_line(frame.as_str().as_bytes())
239    }
240
241    fn decode(
242        &mut self,
243        event: Self::Event,
244        out: Out<'id, Embedding>,
245    ) -> Result<Flow, ProviderError> {
246        if event.usage.is_none() && self.requires_usage {
247            return Err(ProviderError::Response(format!(
248                "{} embedding response omitted required usage",
249                self.provider
250            )));
251        }
252        let usage = event
253            .usage
254            .as_ref()
255            .map(Usage::to_normalized)
256            .unwrap_or_default();
257        // `document` is joined on by the operation's fold, which is the only
258        // place that still holds the request's inputs.
259        let embeddings = event
260            .data
261            .into_iter()
262            .map(|datum| embeddings::Embedding {
263                document: String::new(),
264                vec: datum
265                    .embedding
266                    .into_iter()
267                    .filter_map(|number| number.as_f64())
268                    .collect(),
269            })
270            .collect();
271        let model = if event.model.is_empty() {
272            self.model.clone()
273        } else {
274            event.model
275        };
276        Ok(out.end(embeddings::EmbeddingResponse {
277            model: Some(model),
278            usage,
279            ..embeddings::EmbeddingResponse::new(embeddings)
280        }))
281    }
282}
283
284impl Wire for Embeddings {
285    type Op = Embedding;
286    type Payload = crate::wire::Encoded;
287    type Frame = crate::wire::WireFrame;
288    type Decoder<'id> = EmbeddingsDecoder;
289
290    /// Only an explicit nonzero declaration permits reply-width mismatch checks.
291    fn describe(&self) -> Descriptor<'_> {
292        Descriptor::new(self.provider.dialect.name)
293            .model(self.model.as_str())
294            .capabilities(
295                Capabilities::embedding(
296                    self.provider.dialect.quirks.embedding.max_documents,
297                    self.resolved_ndims(),
298                )
299                .declaring(self.ndims),
300            )
301    }
302
303    fn encode(&self, request: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
304        let quirks = &self.provider.dialect.quirks.embedding;
305        // Base64 vectors are not decoded anywhere, so asking for them would
306        // answer 200 with a payload rig cannot read.
307        if self.encoding_format == Some(EncodingFormat::Base64) {
308            return Err(EncodeError::request(format!(
309                "Rig cannot decode {} embedding responses encoded as `base64`",
310                self.provider.dialect.name
311            )));
312        }
313        if self.encoding_format.is_some() && !quirks.supports_encoding_format {
314            return Err(unsupported_parameter(
315                self.provider.dialect.name,
316                "encoding_format",
317            ));
318        }
319        if self.user.is_some() && !quirks.supports_user {
320            return Err(unsupported_parameter(self.provider.dialect.name, "user"));
321        }
322        self.refuse_unhonourable_width()?;
323
324        let mut body = serde_json::json!({ "input": request });
325        let Some(object) = body.as_object_mut() else {
326            return Err(EncodeError::request(
327                "embedding request body must be an object",
328            ));
329        };
330        if quirks.sends_model_field {
331            object.insert("model".to_owned(), serde_json::json!(self.model));
332        }
333        if let Some((field, ndims)) = self.requested_width() {
334            object.insert(field.to_owned(), serde_json::json!(ndims));
335        }
336        if let Some(encoding_format) = self.encoding_format {
337            object.insert(
338                "encoding_format".to_owned(),
339                serde_json::to_value(encoding_format)?,
340            );
341        }
342        if let Some(user) = &self.user {
343            object.insert("user".to_owned(), serde_json::json!(user));
344        }
345
346        json_post(
347            &self.provider,
348            self.provider.dialect.quirks.embeddings_path,
349            self.provider.deployment(&self.model),
350            &body,
351        )
352    }
353
354    fn decoder<'id>(&self) -> Self::Decoder<'id> {
355        EmbeddingsDecoder {
356            requires_usage: self.provider.dialect.quirks.embedding.requires_usage,
357            provider: self.provider.dialect.name,
358            model: self.model.clone(),
359        }
360    }
361}
362
363/// Transcription requests encoded as multipart or input-audio JSON.
364#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
365pub struct Transcriptions {
366    /// Which provider, and how to reach it.
367    pub provider: OpenAIConfig,
368    /// Transcription model or Azure deployment identifier.
369    pub model: String,
370}
371
372impl Transcriptions {
373    /// The transcription wire for `model`.
374    pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
375        Self {
376            provider,
377            model: model.into(),
378        }
379    }
380
381    /// OpenAI's multipart upload: the audio as a file part beside the
382    /// per-request options.
383    fn multipart_body(&self, request: TranscriptionRequest) -> Result<Body, EncodeError> {
384        use crate::http_client::MultipartForm;
385        use crate::http_client::multipart::Part;
386
387        let mut form = MultipartForm::new();
388        // Azure addresses its deployment through the URL rather than the form.
389        if self.provider.deployment(&self.model).is_none() {
390            form = form.text("model", self.model.clone());
391        }
392        form = form.part(Part::bytes("file", request.data).filename(request.filename));
393        if let Some(language) = request.language {
394            form = form.text("language", language);
395        }
396        if let Some(prompt) = request.prompt {
397            form = form.text("prompt", prompt);
398        }
399        if let Some(temperature) = request.temperature {
400            form = form.text("temperature", temperature.to_string());
401        }
402        if let Some(additional_params) = request.additional_params {
403            for (name, value) in additional_params_object(&additional_params)? {
404                // Form strings must not acquire JSON quotation marks.
405                let value = match value {
406                    serde_json::Value::String(value) => value.clone(),
407                    other => other.to_string(),
408                };
409                form = form.text(name.clone(), value);
410            }
411        }
412        Ok(Body::Multipart(form))
413    }
414
415    /// Encode base64 audio under `input_audio`, inferring format from the filename.
416    /// Reject top-level prompts and non-object additional parameters.
417    fn input_audio_body(&self, request: TranscriptionRequest) -> Result<Body, EncodeError> {
418        use base64::Engine;
419
420        if request.prompt.is_some() {
421            return Err(EncodeError::request(std::io::Error::new(
422                std::io::ErrorKind::InvalidInput,
423                "OpenRouter STT does not support a top-level prompt field. \
424                     Provider-specific prompt options can be passed via `additional_params`. \
425                     Example: {\"provider\": {\"options\": {\"<provider>\": {\"prompt\": \"<text>\"}}}}",
426            )));
427        }
428
429        let mut body = serde_json::Map::new();
430        body.insert("model".to_owned(), serde_json::json!(self.model));
431        body.insert(
432            "input_audio".to_owned(),
433            serde_json::json!({
434                "data": base64::engine::general_purpose::STANDARD.encode(&request.data),
435                "format": audio_format_of(&request.filename),
436            }),
437        );
438        if let Some(language) = request.language {
439            body.insert("language".to_owned(), serde_json::json!(language));
440        }
441        if let Some(temperature) = request.temperature {
442            body.insert("temperature".to_owned(), serde_json::json!(temperature));
443        }
444        if let Some(additional_params) = request.additional_params {
445            for (name, value) in additional_params_object(&additional_params)? {
446                body.insert(name.clone(), value.clone());
447            }
448        }
449        Ok(Body::Bytes(serde_json::to_vec(
450            &serde_json::Value::Object(body),
451        )?))
452    }
453}
454
455/// A transcription request's `additional_params`, as an object.
456fn additional_params_object(
457    params: &serde_json::Value,
458) -> Result<&serde_json::Map<String, serde_json::Value>, EncodeError> {
459    params.as_object().ok_or_else(|| {
460        EncodeError::request(std::io::Error::new(
461            std::io::ErrorKind::InvalidInput,
462            "additional transcription parameters must be a JSON object",
463        ))
464    })
465}
466
467/// Infer the audio container from a case-insensitive filename extension.
468/// Unknown or absent extensions default to `wav`.
469fn audio_format_of(filename: &str) -> &'static str {
470    let extension = std::path::Path::new(filename)
471        .extension()
472        .and_then(std::ffi::OsStr::to_str)
473        .map(str::to_ascii_lowercase);
474    match extension.as_deref() {
475        Some("mp3") => "mp3",
476        Some("flac") => "flac",
477        Some("m4a") => "m4a",
478        Some("ogg") => "ogg",
479        Some("webm") => "webm",
480        Some("aac") => "aac",
481        _ => "wav",
482    }
483}
484
485/// The transcription decoder.
486#[derive(Default)]
487pub struct TranscriptionsDecoder;
488
489impl<'id> Decoder<'id, Transcription> for TranscriptionsDecoder {
490    type Event = crate::providers::openai::transcription::TranscriptionResponse;
491
492    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
493        classify_untyped_line(frame.as_str().as_bytes())
494    }
495
496    fn decode(
497        &mut self,
498        event: Self::Event,
499        out: Out<'id, Transcription>,
500    ) -> Result<Flow, ProviderError> {
501        use crate::transcription::NormalizeTranscriptionResponse;
502        Ok(out.end(event.normalize()?))
503    }
504}
505
506impl Wire for Transcriptions {
507    type Op = Transcription;
508    type Payload = crate::wire::Encoded;
509    type Frame = crate::wire::WireFrame;
510    type Decoder<'id> = TranscriptionsDecoder;
511
512    fn describe(&self) -> Descriptor<'_> {
513        Descriptor::new(self.provider.dialect.name).model(self.model.as_str())
514    }
515
516    fn encode(&self, request: TranscriptionRequest, _mode: Mode) -> Result<Encoded, EncodeError> {
517        let uri = self
518            .provider
519            .modality_uri(
520                "transcription",
521                self.provider.dialect.quirks.transcription_path,
522                &self.model,
523            )
524            .map_err(EncodeError::request)?;
525        let builder = http::Request::post(uri);
526        let (builder, body) = match self.provider.dialect.quirks.transcription_body {
527            TranscriptionBody::Multipart => (builder, self.multipart_body(request)?),
528            TranscriptionBody::InputAudioJson => (
529                builder.header(http::header::CONTENT_TYPE, "application/json"),
530                self.input_audio_body(request)?,
531            ),
532        };
533        encoded(&self.provider, builder, body)
534    }
535
536    fn decoder<'id>(&self) -> Self::Decoder<'id> {
537        TranscriptionsDecoder
538    }
539}
540
541/// The image-generation wire.
542#[cfg(feature = "image")]
543#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
544pub struct Images {
545    /// Which provider, and how to reach it.
546    pub provider: OpenAIConfig,
547    /// The image model.
548    pub model: String,
549}
550
551#[cfg(feature = "image")]
552impl Images {
553    /// The image-generation wire for `model`.
554    pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
555        Self {
556            provider,
557            model: model.into(),
558        }
559    }
560}
561
562/// The image-generation decoder.
563///
564/// Carries the dialect's [`ImageBody`] because the request shape names the
565/// reply shape: most of this family answers with a JSON envelope, and
566/// Hugging Face's router answers with the image bytes themselves.
567#[cfg(feature = "image")]
568#[derive(Default)]
569pub struct ImagesDecoder {
570    /// Which reply shape this dialect answers with.
571    body: ImageBody,
572}
573
574/// One generated image, as every dialect on this wire returns it.
575#[cfg(feature = "image")]
576#[derive(Debug, Clone, Serialize, Deserialize)]
577pub struct ImageDatum {
578    /// The image, base64-encoded.
579    pub b64_json: String,
580}
581
582/// Base64 image represented as a string or an object with an `image` field.
583#[cfg(feature = "image")]
584#[derive(Debug, Clone, Serialize, Deserialize)]
585#[serde(untagged)]
586pub enum ImagesReplyImage {
587    /// Hyperbolic: `{"image": "<base64>"}`.
588    Keyed {
589        /// The image, base64-encoded.
590        image: String,
591    },
592    /// Venice: the base64 payload itself.
593    Bare(String),
594}
595
596#[cfg(feature = "image")]
597impl ImagesReplyImage {
598    /// The image's base64 payload, whichever form the dialect sent.
599    pub fn base64(&self) -> &str {
600        match self {
601            Self::Keyed { image } => image,
602            Self::Bare(image) => image,
603        }
604    }
605}
606
607/// JSON image reply accepting `data` and `images` arrays with optional metadata.
608#[cfg(feature = "image")]
609#[derive(Debug, Clone, Serialize, Deserialize)]
610pub struct ImagesReply {
611    /// The generated images, as OpenAI and xAI key them.
612    #[serde(default)]
613    pub data: Vec<ImageDatum>,
614    /// The generated images, as Hyperbolic and Venice key them.
615    #[serde(default)]
616    pub images: Vec<ImagesReplyImage>,
617    /// Whatever else the dialect sent (`created`, Venice's `id`/`timing`, and
618    /// any field this build does not model), so the raw payload loses
619    /// nothing.
620    #[serde(flatten)]
621    pub extra: serde_json::Map<String, serde_json::Value>,
622}
623
624#[cfg(feature = "image")]
625impl ImagesReply {
626    /// The first image's base64 payload, whichever key the dialect used.
627    pub fn first_base64(&self) -> Option<&str> {
628        self.data
629            .first()
630            .map(|image| image.b64_json.as_str())
631            .or_else(|| self.images.first().map(ImagesReplyImage::base64))
632            .filter(|encoded| !encoded.is_empty())
633    }
634}
635
636/// Generated images in a decoded JSON envelope or as raw bytes.
637#[cfg(feature = "image")]
638#[derive(Debug, Clone)]
639pub enum ImagesEvent {
640    /// A JSON envelope, as OpenAI, xAI and Hyperbolic answer.
641    Json(ImagesReply),
642    /// The image bytes themselves, with no envelope at all.
643    Raw(Vec<u8>),
644}
645
646#[cfg(feature = "image")]
647impl<'id> Decoder<'id, crate::operation::ImageGeneration> for ImagesDecoder {
648    type Event = ImagesEvent;
649
650    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
651        match self.body {
652            // Preserve raw image bytes even when framing recognized valid UTF-8.
653            ImageBody::HuggingFace => WireEvent::Known(ImagesEvent::Raw(match frame {
654                WireFrame::Text(text) => text.into_bytes(),
655                WireFrame::Bytes(bytes) => bytes,
656            })),
657            ImageBody::OpenAi | ImageBody::Xai | ImageBody::Hyperbolic | ImageBody::Venice => {
658                classify_untyped_line(frame.as_str().as_bytes()).map(ImagesEvent::Json)
659            }
660        }
661    }
662
663    fn decode(
664        &mut self,
665        event: Self::Event,
666        out: Out<'id, crate::operation::ImageGeneration>,
667    ) -> Result<Flow, ProviderError> {
668        use crate::image_generation::ImageGenerationResponse;
669        use base64::Engine;
670
671        let reply = match event {
672            // The image is already the payload, and the reply is not a
673            // document, so `raw` stays null rather than restating the bytes.
674            ImagesEvent::Raw(image) => {
675                return Ok(out.end(ImageGenerationResponse::new(image)));
676            }
677            ImagesEvent::Json(reply) => reply,
678        };
679        let Some(encoded) = reply.first_base64() else {
680            return Err(ProviderError::Response("missing image data".to_owned()));
681        };
682        let image = match base64::prelude::BASE64_STANDARD.decode(encoded) {
683            Ok(image) => image,
684            Err(error) => {
685                return Err(ProviderError::Response(error.to_string()));
686            }
687        };
688        Ok(out.end(ImageGenerationResponse::new(image)))
689    }
690}
691
692#[cfg(feature = "image")]
693impl Wire for Images {
694    type Op = crate::operation::ImageGeneration;
695    type Payload = crate::wire::Encoded;
696    type Frame = crate::wire::WireFrame;
697    type Decoder<'id> = ImagesDecoder;
698
699    fn describe(&self) -> Descriptor<'_> {
700        Descriptor::new(self.provider.dialect.name).model(self.model.as_str())
701    }
702
703    fn encode(
704        &self,
705        request: crate::image_generation::ImageGenerationRequest,
706        _mode: Mode,
707    ) -> Result<Encoded, EncodeError> {
708        let mut body = match self.provider.dialect.quirks.image_body {
709            // `response_format` is deliberately absent: it is no longer part
710            // of OpenAI's request schema, which rejects it before it even
711            // looks at the model. A compatible endpoint that still takes the
712            // field gets it through `additional_params`.
713            ImageBody::OpenAi => serde_json::json!({
714                "model": self.model,
715                "prompt": request.prompt,
716                "size": format!("{}x{}", request.width, request.height),
717            }),
718            // xAI takes no `size` and answers with a URL unless asked for
719            // base64, which is the only form this wire decodes.
720            ImageBody::Xai => serde_json::json!({
721                "model": self.model,
722                "prompt": request.prompt,
723                "response_format": "b64_json",
724                "aspect_ratio": "1:1",
725            }),
726            ImageBody::Hyperbolic => serde_json::json!({
727                "model_name": self.model,
728                "prompt": request.prompt,
729                "height": request.height,
730                "width": request.width,
731            }),
732            ImageBody::Venice => serde_json::json!({
733                "model": self.model,
734                "prompt": request.prompt,
735                "width": request.width,
736                "height": request.height,
737            }),
738            // The model is addressed through the URL, not the body.
739            ImageBody::HuggingFace => serde_json::json!({
740                "inputs": request.prompt,
741                "parameters": {
742                    "width": request.width,
743                    "height": request.height,
744                },
745            }),
746        };
747        // Merged last, so a caller can reach the endpoint's other parameters
748        // and override what is derived above.
749        if let Some(additional_params) = request.additional_params {
750            crate::json_utils::merge_inplace(&mut body, additional_params);
751        }
752
753        let uri = self
754            .provider
755            .modality_uri(
756                "image generation",
757                self.provider.dialect.quirks.image_generation_path,
758                &self.model,
759            )
760            .map_err(EncodeError::request)?;
761        json_post_to(&self.provider, uri, &body)
762    }
763
764    fn decoder<'id>(&self) -> Self::Decoder<'id> {
765        ImagesDecoder {
766            body: self.provider.dialect.quirks.image_body,
767        }
768    }
769}
770
771/// The speech wire.
772#[cfg(feature = "audio")]
773#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
774pub struct Speech {
775    /// Which provider, and how to reach it.
776    pub provider: OpenAIConfig,
777    /// The speech model.
778    pub model: String,
779}
780
781#[cfg(feature = "audio")]
782impl Speech {
783    /// The speech wire for `model`.
784    pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
785        Self {
786            provider,
787            model: model.into(),
788        }
789    }
790}
791
792/// Decode raw audio or the dialect's base64 JSON envelope.
793#[cfg(feature = "audio")]
794#[derive(Default)]
795pub struct SpeechDecoder {
796    /// Which reply shape this dialect answers with.
797    body: SpeechBody,
798}
799
800/// Hyperbolic's speech reply: base64 in a JSON envelope rather than the
801/// audio bytes themselves.
802#[cfg(feature = "audio")]
803#[derive(Debug, Clone, Serialize, Deserialize)]
804pub struct SpeechReply {
805    /// The audio, base64-encoded.
806    pub audio: String,
807}
808
809#[cfg(feature = "audio")]
810impl<'id> Decoder<'id, crate::operation::AudioGeneration> for SpeechDecoder {
811    type Event = Vec<u8>;
812
813    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
814        // Interpretation selects raw audio or JSON decoding from the dialect.
815        WireEvent::Known(match frame {
816            WireFrame::Text(text) => text.into_bytes(),
817            WireFrame::Bytes(bytes) => bytes,
818        })
819    }
820
821    fn decode(
822        &mut self,
823        event: Self::Event,
824        out: Out<'id, crate::operation::AudioGeneration>,
825    ) -> Result<Flow, ProviderError> {
826        use base64::Engine;
827
828        let audio = match self.body {
829            // OpenAI and xAI answer with the audio itself.
830            SpeechBody::OpenAi | SpeechBody::Xai => event,
831            // Hyperbolic wraps it, base64-encoded, in a JSON envelope.
832            SpeechBody::Hyperbolic => {
833                let reply = match serde_json::from_slice::<SpeechReply>(&event) {
834                    Ok(reply) => reply,
835                    Err(error) => {
836                        return Err(ProviderError::Response(error.to_string()));
837                    }
838                };
839                match base64::prelude::BASE64_STANDARD.decode(&reply.audio) {
840                    Ok(audio) => audio,
841                    Err(error) => {
842                        return Err(ProviderError::Response(error.to_string()));
843                    }
844                }
845            }
846        };
847        Ok(out.end(crate::audio_generation::AudioGenerationResponse::new(audio)))
848    }
849}
850
851#[cfg(feature = "audio")]
852impl Wire for Speech {
853    type Op = crate::operation::AudioGeneration;
854    type Payload = crate::wire::Encoded;
855    type Frame = crate::wire::WireFrame;
856    type Decoder<'id> = SpeechDecoder;
857
858    fn describe(&self) -> Descriptor<'_> {
859        Descriptor::new(self.provider.dialect.name).model(self.model.as_str())
860    }
861
862    fn encode(
863        &self,
864        request: crate::audio_generation::AudioGenerationRequest,
865        _mode: Mode,
866    ) -> Result<Encoded, EncodeError> {
867        let mut body = match self.provider.dialect.quirks.speech_body {
868            SpeechBody::OpenAi => serde_json::json!({
869                "model": self.model,
870                "input": request.text,
871                "voice": request.voice,
872                "speed": request.speed,
873            }),
874            // xAI requires a default voice when the caller leaves it empty.
875            SpeechBody::Xai => serde_json::json!({
876                "text": request.text,
877                "voice_id": if request.voice.is_empty() { "eve" } else { request.voice.as_str() },
878                "language": "en",
879            }),
880            // Hyperbolic addresses this endpoint by language, so the
881            // identifier the caller passes as the model IS the language tag.
882            SpeechBody::Hyperbolic => serde_json::json!({
883                "language": self.model,
884                "speaker": request.voice,
885                "text": request.text,
886                "speed": request.speed,
887            }),
888        };
889        // Caller parameters take precedence, including response format and instructions.
890        if let Some(additional_params) = request.additional_params {
891            crate::json_utils::merge_inplace(&mut body, additional_params);
892        }
893
894        // Azure versions its speech endpoint separately from every other
895        // route, so this one request carries its own `api-version`.
896        let uri = self.provider.uri_versioned(
897            self.provider.dialect.quirks.audio_generation_path,
898            self.provider.deployment(&self.model),
899            self.provider.speech_api_version(),
900        );
901        json_post_to(&self.provider, uri, &body)
902    }
903
904    fn decoder<'id>(&self) -> Self::Decoder<'id> {
905        SpeechDecoder {
906            body: self.provider.dialect.quirks.speech_body,
907        }
908    }
909}
910
911/// The model-listing wire.
912#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
913pub struct Models {
914    /// Which provider, and how to reach it.
915    pub provider: OpenAIConfig,
916}
917
918impl Models {
919    /// The model-listing wire.
920    pub fn new(provider: OpenAIConfig) -> Self {
921        Self { provider }
922    }
923}
924
925/// Model-list entry with required `id` and optional metadata.
926/// Context length prefers `context_window`, then `context_length`, then
927/// `max_context_length`. Top-level output limits take precedence over `top_provider`.
928#[derive(Debug, Deserialize)]
929pub struct ModelEntry {
930    pub id: String,
931    #[serde(default)]
932    pub name: Option<String>,
933    #[serde(default)]
934    pub description: Option<String>,
935    /// Mistral labels the model kind `type` (`base`, `fine-tuned`).
936    #[serde(default, rename = "type")]
937    pub kind: Option<String>,
938    #[serde(default)]
939    pub created: Option<u64>,
940    #[serde(default)]
941    pub owned_by: Option<String>,
942    #[serde(default)]
943    pub context_window: Option<u32>,
944    #[serde(default)]
945    pub context_length: Option<u32>,
946    #[serde(default)]
947    pub max_context_length: Option<u32>,
948    #[serde(default)]
949    pub max_completion_tokens: Option<u32>,
950    #[serde(default)]
951    pub top_provider: Option<TopProvider>,
952}
953
954/// OpenRouter's per-entry routing block. Only the output ceiling is read;
955/// the rest of the block is routing detail [`ModelInfo`] has no slot for.
956#[derive(Debug, Deserialize)]
957pub struct TopProvider {
958    #[serde(default)]
959    pub max_completion_tokens: Option<u32>,
960}
961
962impl From<ModelEntry> for ModelInfo {
963    fn from(entry: ModelEntry) -> Self {
964        let mut model = ModelInfo::from_id(entry.id);
965        model.name = entry.name;
966        model.description = entry.description;
967        model.r#type = entry.kind;
968        model.created_at = entry.created;
969        model.owned_by = entry.owned_by;
970        model.context_length = entry
971            .context_window
972            .or(entry.context_length)
973            .or(entry.max_context_length);
974        model.max_output_tokens = entry.max_completion_tokens.or_else(|| {
975            entry
976                .top_provider
977                .and_then(|provider| provider.max_completion_tokens)
978        });
979        model
980    }
981}
982
983/// The `{ "data": [...] }` envelope.
984#[derive(Debug, Deserialize)]
985pub struct ModelsReply {
986    #[serde(default)]
987    pub data: Vec<ModelEntry>,
988}
989
990/// Decode a complete model catalogue without pagination.
991#[derive(Default)]
992pub struct ModelsDecoder;
993
994impl<'id> Decoder<'id, ModelListing> for ModelsDecoder {
995    type Event = ModelsReply;
996
997    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
998        classify_untyped_line(frame.as_str().as_bytes())
999    }
1000
1001    fn decode(
1002        &mut self,
1003        event: Self::Event,
1004        out: Out<'id, ModelListing>,
1005    ) -> Result<Flow, ProviderError> {
1006        let models = event.data.into_iter().map(ModelInfo::from).collect();
1007        Ok(out.end(ModelPage {
1008            models: ModelList::new(models),
1009            next: None,
1010        }))
1011    }
1012}
1013
1014impl Wire for Models {
1015    type Op = ModelListing;
1016    type Payload = crate::wire::Encoded;
1017    type Frame = crate::wire::WireFrame;
1018    type Decoder<'id> = ModelsDecoder;
1019
1020    fn describe(&self) -> Descriptor<'_> {
1021        Descriptor::new(self.provider.dialect.name)
1022    }
1023
1024    /// The catalogue is one page, so a cursor never follows.
1025    fn encode(&self, _cursor: Option<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
1026        get(&self.provider, self.provider.dialect.quirks.models_path)
1027    }
1028
1029    fn decoder<'id>(&self) -> Self::Decoder<'id> {
1030        ModelsDecoder
1031    }
1032}
1033
1034/// Rerank documents with `{model, query, documents, top_n}` requests.
1035/// The dialect must configure a nonempty reranking path.
1036#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1037pub struct Rerank {
1038    /// Which provider, and how to reach it.
1039    pub provider: OpenAIConfig,
1040    /// The reranker model.
1041    pub model: String,
1042    /// Return only the `top_n` highest-scoring documents, when the caller
1043    /// asked for a cut.
1044    pub top_n: Option<usize>,
1045}
1046
1047impl Rerank {
1048    /// The rerank wire for `model`.
1049    pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
1050        Self {
1051            provider,
1052            model: model.into(),
1053            top_n: None,
1054        }
1055    }
1056
1057    /// Ask the server to return only the `top_n` highest-scoring documents.
1058    pub fn with_top_n(mut self, top_n: usize) -> Self {
1059        self.top_n = Some(top_n);
1060        self
1061    }
1062}
1063
1064/// Scored input document accepting `relevance_score` or `score` as its score key.
1065#[derive(Debug, Clone, Serialize, Deserialize)]
1066pub struct RerankResultEntry {
1067    /// Which input document this scored.
1068    pub index: usize,
1069    /// The score.
1070    #[serde(alias = "score")]
1071    pub relevance_score: f64,
1072    /// Present only on servers that echo the document back; llama.cpp does
1073    /// not on this path.
1074    #[serde(default, alias = "text")]
1075    pub document: Option<String>,
1076}
1077
1078/// What a rerank reply reports besides its ranking.
1079#[derive(Debug, Clone, Serialize, Deserialize, Default)]
1080pub struct RerankUsage {
1081    /// Tokens the query and documents cost.
1082    #[serde(default)]
1083    pub prompt_tokens: u64,
1084    /// Total tokens, as the provider reported them.
1085    #[serde(default)]
1086    pub total_tokens: u64,
1087}
1088
1089/// The rerank reply.
1090#[derive(Debug, Clone, Serialize, Deserialize)]
1091pub struct RerankReply {
1092    /// The model the server ranked with, when it named one.
1093    #[serde(default)]
1094    pub model: Option<String>,
1095    /// The ranking.
1096    pub results: Vec<RerankResultEntry>,
1097    /// What it cost.
1098    #[serde(default)]
1099    pub usage: Option<RerankUsage>,
1100}
1101
1102/// The rerank decoder.
1103#[derive(Default)]
1104pub struct RerankDecoder;
1105
1106impl<'id> Decoder<'id, RerankOp> for RerankDecoder {
1107    type Event = RerankReply;
1108
1109    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
1110        classify_untyped_line(frame.as_str().as_bytes())
1111    }
1112
1113    fn decode(
1114        &mut self,
1115        event: Self::Event,
1116        out: Out<'id, RerankOp>,
1117    ) -> Result<Flow, ProviderError> {
1118        let usage = event
1119            .usage
1120            .map(|usage| crate::completion::Usage {
1121                input_tokens: Some(usage.prompt_tokens),
1122                total_tokens: Some(usage.total_tokens),
1123                ..Default::default()
1124            })
1125            .unwrap_or_default();
1126        let results = event
1127            .results
1128            .into_iter()
1129            .map(|result| crate::rerank::RerankResult {
1130                index: result.index,
1131                document: result.document,
1132                relevance_score: result.relevance_score,
1133            })
1134            .collect();
1135        Ok(out.end(crate::rerank::RerankResponse {
1136            model: event.model,
1137            usage,
1138            ..crate::rerank::RerankResponse::new(results)
1139        }))
1140    }
1141}
1142
1143impl Wire for Rerank {
1144    type Op = RerankOp;
1145    type Payload = crate::wire::Encoded;
1146    type Frame = crate::wire::WireFrame;
1147    type Decoder<'id> = RerankDecoder;
1148
1149    fn describe(&self) -> Descriptor<'_> {
1150        Descriptor::new(self.provider.dialect.name)
1151            .model(self.model.as_str())
1152            .capabilities(Capabilities::rerank(
1153                self.provider.dialect.quirks.rerank.max_documents,
1154            ))
1155    }
1156
1157    fn encode(
1158        &self,
1159        request: crate::operation::RerankRequest,
1160        _mode: Mode,
1161    ) -> Result<Encoded, EncodeError> {
1162        let quirks = &self.provider.dialect.quirks.rerank;
1163        // An empty path explicitly disables reranking.
1164        if quirks.path.is_empty() {
1165            return Err(EncodeError::request(format!(
1166                "{} offers no reranking endpoint",
1167                self.provider.dialect.name
1168            )));
1169        }
1170        let mut body = serde_json::json!({
1171            "query": request.query,
1172            "documents": request.documents,
1173        });
1174        let Some(object) = body.as_object_mut() else {
1175            return Err(EncodeError::request(
1176                "rerank request body must be an object",
1177            ));
1178        };
1179        if quirks.sends_model_field {
1180            object.insert("model".to_owned(), serde_json::json!(self.model));
1181        }
1182        if let Some(top_n) = self.top_n {
1183            object.insert("top_n".to_owned(), serde_json::json!(top_n));
1184        }
1185
1186        json_post(
1187            &self.provider,
1188            quirks.path,
1189            self.provider.deployment(&self.model),
1190            &body,
1191        )
1192    }
1193
1194    fn decoder<'id>(&self) -> Self::Decoder<'id> {
1195        RerankDecoder
1196    }
1197}
1198
1199/// The credential-check wire: a `GET` whose status is the answer.
1200#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1201pub struct Verify {
1202    /// Which provider, and how to reach it.
1203    pub provider: OpenAIConfig,
1204}
1205
1206impl Verify {
1207    /// The credential-check wire.
1208    pub fn new(provider: OpenAIConfig) -> Self {
1209        Self { provider }
1210    }
1211}
1212
1213pub use crate::operation::VerifyDecoder;
1214
1215impl Wire for Verify {
1216    type Op = VerifyOp;
1217    type Payload = crate::wire::Encoded;
1218    type Frame = crate::wire::WireFrame;
1219    type Decoder<'id> = VerifyDecoder;
1220
1221    fn describe(&self) -> Descriptor<'_> {
1222        Descriptor::new(self.provider.dialect.name)
1223    }
1224
1225    fn encode(&self, _request: (), _mode: Mode) -> Result<Encoded, EncodeError> {
1226        let path = self.provider.dialect.quirks.verify_path;
1227        if path.is_empty() {
1228            return Err(EncodeError::request(format!(
1229                "{} offers no endpoint that checks a credential without consuming tokens",
1230                self.provider.dialect.name
1231            )));
1232        }
1233        get(&self.provider, path)
1234    }
1235
1236    fn decoder<'id>(&self) -> Self::Decoder<'id> {
1237        VerifyDecoder
1238    }
1239}
1240
1241#[cfg(test)]
1242mod tests;