1use 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
36fn 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
46fn 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
58fn 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
68fn 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
89fn unsupported_parameter(provider: &str, parameter: &str) -> EncodeError {
91 EncodeError::request(format!(
92 "{provider} embeddings do not support the `{parameter}` parameter"
93 ))
94}
95
96#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
98pub struct Embeddings {
99 pub provider: OpenAIConfig,
101 pub model: String,
103 pub ndims: Option<usize>,
106 pub encoding_format: Option<EncodingFormat>,
108 pub user: Option<String>,
110}
111
112impl Embeddings {
113 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 pub fn with_encoding_format(mut self, encoding_format: EncodingFormat) -> Self {
126 self.encoding_format = Some(encoding_format);
127 self
128 }
129
130 pub fn with_user(mut self, user: impl Into<String>) -> Self {
132 self.user = Some(user.into());
133 self
134 }
135
136 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 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 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 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 let Some(width) = self.model_width() else {
185 return Ok(());
186 };
187 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 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 let ndims = match self.resolved_ndims() {
209 0 => return None,
210 ndims => ndims,
211 };
212 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#[derive(Default)]
227pub struct EmbeddingsDecoder {
228 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 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 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 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
365pub struct Transcriptions {
366 pub provider: OpenAIConfig,
368 pub model: String,
370}
371
372impl Transcriptions {
373 pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
375 Self {
376 provider,
377 model: model.into(),
378 }
379 }
380
381 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 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 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 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
455fn 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
467fn 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#[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#[cfg(feature = "image")]
543#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
544pub struct Images {
545 pub provider: OpenAIConfig,
547 pub model: String,
549}
550
551#[cfg(feature = "image")]
552impl Images {
553 pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
555 Self {
556 provider,
557 model: model.into(),
558 }
559 }
560}
561
562#[cfg(feature = "image")]
568#[derive(Default)]
569pub struct ImagesDecoder {
570 body: ImageBody,
572}
573
574#[cfg(feature = "image")]
576#[derive(Debug, Clone, Serialize, Deserialize)]
577pub struct ImageDatum {
578 pub b64_json: String,
580}
581
582#[cfg(feature = "image")]
584#[derive(Debug, Clone, Serialize, Deserialize)]
585#[serde(untagged)]
586pub enum ImagesReplyImage {
587 Keyed {
589 image: String,
591 },
592 Bare(String),
594}
595
596#[cfg(feature = "image")]
597impl ImagesReplyImage {
598 pub fn base64(&self) -> &str {
600 match self {
601 Self::Keyed { image } => image,
602 Self::Bare(image) => image,
603 }
604 }
605}
606
607#[cfg(feature = "image")]
609#[derive(Debug, Clone, Serialize, Deserialize)]
610pub struct ImagesReply {
611 #[serde(default)]
613 pub data: Vec<ImageDatum>,
614 #[serde(default)]
616 pub images: Vec<ImagesReplyImage>,
617 #[serde(flatten)]
621 pub extra: serde_json::Map<String, serde_json::Value>,
622}
623
624#[cfg(feature = "image")]
625impl ImagesReply {
626 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#[cfg(feature = "image")]
638#[derive(Debug, Clone)]
639pub enum ImagesEvent {
640 Json(ImagesReply),
642 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 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 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 ImageBody::OpenAi => serde_json::json!({
714 "model": self.model,
715 "prompt": request.prompt,
716 "size": format!("{}x{}", request.width, request.height),
717 }),
718 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 ImageBody::HuggingFace => serde_json::json!({
740 "inputs": request.prompt,
741 "parameters": {
742 "width": request.width,
743 "height": request.height,
744 },
745 }),
746 };
747 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#[cfg(feature = "audio")]
773#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
774pub struct Speech {
775 pub provider: OpenAIConfig,
777 pub model: String,
779}
780
781#[cfg(feature = "audio")]
782impl Speech {
783 pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
785 Self {
786 provider,
787 model: model.into(),
788 }
789 }
790}
791
792#[cfg(feature = "audio")]
794#[derive(Default)]
795pub struct SpeechDecoder {
796 body: SpeechBody,
798}
799
800#[cfg(feature = "audio")]
803#[derive(Debug, Clone, Serialize, Deserialize)]
804pub struct SpeechReply {
805 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 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 SpeechBody::OpenAi | SpeechBody::Xai => event,
831 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 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 SpeechBody::Hyperbolic => serde_json::json!({
883 "language": self.model,
884 "speaker": request.voice,
885 "text": request.text,
886 "speed": request.speed,
887 }),
888 };
889 if let Some(additional_params) = request.additional_params {
891 crate::json_utils::merge_inplace(&mut body, additional_params);
892 }
893
894 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
913pub struct Models {
914 pub provider: OpenAIConfig,
916}
917
918impl Models {
919 pub fn new(provider: OpenAIConfig) -> Self {
921 Self { provider }
922 }
923}
924
925#[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 #[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#[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#[derive(Debug, Deserialize)]
985pub struct ModelsReply {
986 #[serde(default)]
987 pub data: Vec<ModelEntry>,
988}
989
990#[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 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1037pub struct Rerank {
1038 pub provider: OpenAIConfig,
1040 pub model: String,
1042 pub top_n: Option<usize>,
1045}
1046
1047impl Rerank {
1048 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 pub fn with_top_n(mut self, top_n: usize) -> Self {
1059 self.top_n = Some(top_n);
1060 self
1061 }
1062}
1063
1064#[derive(Debug, Clone, Serialize, Deserialize)]
1066pub struct RerankResultEntry {
1067 pub index: usize,
1069 #[serde(alias = "score")]
1071 pub relevance_score: f64,
1072 #[serde(default, alias = "text")]
1075 pub document: Option<String>,
1076}
1077
1078#[derive(Debug, Clone, Serialize, Deserialize, Default)]
1080pub struct RerankUsage {
1081 #[serde(default)]
1083 pub prompt_tokens: u64,
1084 #[serde(default)]
1086 pub total_tokens: u64,
1087}
1088
1089#[derive(Debug, Clone, Serialize, Deserialize)]
1091pub struct RerankReply {
1092 #[serde(default)]
1094 pub model: Option<String>,
1095 pub results: Vec<RerankResultEntry>,
1097 #[serde(default)]
1099 pub usage: Option<RerankUsage>,
1100}
1101
1102#[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 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1201pub struct Verify {
1202 pub provider: OpenAIConfig,
1204}
1205
1206impl Verify {
1207 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;