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::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#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
34pub enum ImageBody {
35 #[default]
37 OpenAi,
38 Xai,
41 Hyperbolic,
43 Venice,
45 HuggingFace,
48}
49
50#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
52pub enum SpeechBody {
53 #[default]
55 OpenAi,
56 Xai,
58 Hyperbolic,
64}
65
66#[derive(Clone, Copy, Debug, PartialEq, Eq)]
68pub enum TranscriptionBody {
69 Multipart,
72 InputAudioJson,
78}
79
80#[derive(Clone, Copy, Debug, PartialEq, Eq)]
82pub enum DimensionsField {
83 Dimensions,
85 OutputDimension,
87 Ignored,
90}
91
92impl DimensionsField {
93 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
109pub enum AcceptedWidths {
110 Fixed,
114 Range {
117 min: usize,
119 max: usize,
121 requirement: &'static str,
124 },
125}
126
127#[derive(Clone, Copy, Debug, PartialEq, Eq)]
133#[non_exhaustive]
134pub struct ModelWidth {
135 pub model: &'static str,
137 pub default: Option<usize>,
140 pub accepted: AcceptedWidths,
142}
143
144#[derive(Clone, Copy, Debug, PartialEq, Eq)]
146#[non_exhaustive]
147pub struct RerankQuirks {
148 pub path: &'static str,
150 pub max_documents: usize,
152 pub sends_model_field: bool,
154}
155
156impl RerankQuirks {
157 pub const fn unsupported() -> Self {
159 Self {
160 path: "",
161 max_documents: 0,
162 sends_model_field: true,
163 }
164 }
165}
166
167#[derive(Clone, Copy, Debug, PartialEq, Eq)]
169#[non_exhaustive]
170pub struct EmbeddingQuirks {
171 pub max_documents: usize,
173 pub requires_usage: bool,
175 pub supports_encoding_format: bool,
177 pub supports_user: bool,
179 pub sends_model_field: bool,
182 pub dimensions: DimensionsField,
184 pub widths: &'static [ModelWidth],
187 pub refuse_zero_width: Option<&'static str>,
191}
192
193impl EmbeddingQuirks {
194 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 widths: &[],
205 refuse_zero_width: None,
207 }
208 }
209}
210
211impl super::SubRoute {
212 pub fn serves_model_routed_endpoints(&self) -> bool {
215 matches!(self, Self::HFInference)
216 }
217}
218
219impl OpenAIConfig {
220 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 pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
228 Embeddings::new(self.clone(), model, ndims)
229 }
230
231 pub(crate) fn rerank(&self, model: impl Into<String>) -> Rerank {
233 Rerank::new(self.clone(), model)
234 }
235
236 pub(crate) fn transcription(&self, model: impl Into<String>) -> Transcriptions {
238 Transcriptions::new(self.clone(), model)
239 }
240
241 pub(crate) fn models(&self) -> Models {
243 Models::new(self.clone())
244 }
245
246 pub(crate) fn verify(&self) -> Verify {
248 Verify::new(self.clone())
249 }
250
251 #[cfg(feature = "image")]
253 pub(crate) fn image_generation(&self, model: impl Into<String>) -> Images {
254 Images::new(self.clone(), model)
255 }
256
257 #[cfg(feature = "audio")]
259 pub(crate) fn audio_generation(&self, model: impl Into<String>) -> Speech {
260 Speech::new(self.clone(), model)
261 }
262
263 #[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 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
296pub(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 default: None,
322 accepted: AcceptedWidths::Range {
323 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
341pub(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
354fn 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
364fn 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
376fn 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
386fn 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
407fn unsupported_parameter(provider: &str, parameter: &str) -> EncodeError {
409 EncodeError::request(format!(
410 "{provider} embeddings do not support the `{parameter}` parameter"
411 ))
412}
413
414#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
416pub struct Embeddings {
417 pub provider: OpenAIConfig,
419 pub model: String,
421 pub ndims: Option<usize>,
424 pub encoding_format: Option<EncodingFormat>,
426 pub user: Option<String>,
428}
429
430impl Embeddings {
431 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 pub fn with_encoding_format(mut self, encoding_format: EncodingFormat) -> Self {
444 self.encoding_format = Some(encoding_format);
445 self
446 }
447
448 pub fn with_user(mut self, user: impl Into<String>) -> Self {
450 self.user = Some(user.into());
451 self
452 }
453
454 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 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 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 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 let Some(width) = self.model_width() else {
503 return Ok(());
504 };
505 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 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 let ndims = match self.resolved_ndims() {
527 0 => return None,
528 ndims => ndims,
529 };
530 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#[derive(Default)]
545pub struct EmbeddingsDecoder {
546 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 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 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
677pub struct Transcriptions {
678 pub provider: OpenAIConfig,
680 pub model: String,
682}
683
684impl Transcriptions {
685 pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
687 Self {
688 provider,
689 model: model.into(),
690 }
691 }
692
693 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 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 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 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
767fn 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
779fn 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#[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#[cfg(feature = "image")]
855#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
856pub struct Images {
857 pub provider: OpenAIConfig,
859 pub model: String,
861}
862
863#[cfg(feature = "image")]
864impl Images {
865 pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
867 Self {
868 provider,
869 model: model.into(),
870 }
871 }
872}
873
874#[cfg(feature = "image")]
880#[derive(Default)]
881pub struct ImagesDecoder {
882 body: ImageBody,
884}
885
886#[cfg(feature = "image")]
888#[derive(Debug, Clone, Serialize, Deserialize)]
889pub struct ImageDatum {
890 pub b64_json: String,
892}
893
894#[cfg(feature = "image")]
896#[derive(Debug, Clone, Serialize, Deserialize)]
897#[serde(untagged)]
898pub enum ImagesReplyImage {
899 Keyed {
901 image: String,
903 },
904 Bare(String),
906}
907
908#[cfg(feature = "image")]
909impl ImagesReplyImage {
910 pub fn base64(&self) -> &str {
912 match self {
913 Self::Keyed { image } => image,
914 Self::Bare(image) => image,
915 }
916 }
917}
918
919#[cfg(feature = "image")]
921#[derive(Debug, Clone, Serialize, Deserialize)]
922pub struct ImagesReply {
923 #[serde(default)]
925 pub data: Vec<ImageDatum>,
926 #[serde(default)]
928 pub images: Vec<ImagesReplyImage>,
929 #[serde(flatten)]
933 pub extra: serde_json::Map<String, serde_json::Value>,
934}
935
936#[cfg(feature = "image")]
937impl ImagesReply {
938 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#[cfg(feature = "image")]
950#[derive(Debug, Clone)]
951pub enum ImagesEvent {
952 Json(ImagesReply),
954 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 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 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 ImageBody::OpenAi => serde_json::json!({
1027 "model": self.model,
1028 "prompt": request.prompt,
1029 "size": format!("{}x{}", request.width, request.height),
1030 }),
1031 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 ImageBody::HuggingFace => serde_json::json!({
1053 "inputs": request.prompt,
1054 "parameters": {
1055 "width": request.width,
1056 "height": request.height,
1057 },
1058 }),
1059 };
1060 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#[cfg(feature = "audio")]
1086#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1087pub struct Speech {
1088 pub provider: OpenAIConfig,
1090 pub model: String,
1092}
1093
1094#[cfg(feature = "audio")]
1095impl Speech {
1096 pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
1098 Self {
1099 provider,
1100 model: model.into(),
1101 }
1102 }
1103}
1104
1105#[cfg(feature = "audio")]
1107#[derive(Default)]
1108pub struct SpeechDecoder {
1109 body: SpeechBody,
1111}
1112
1113#[cfg(feature = "audio")]
1116#[derive(Debug, Clone, Serialize, Deserialize)]
1117pub struct SpeechReply {
1118 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 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 SpeechBody::OpenAi | SpeechBody::Xai => event,
1144 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 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 SpeechBody::Hyperbolic => serde_json::json!({
1197 "language": self.model,
1198 "speaker": request.voice,
1199 "text": request.text,
1200 "speed": request.speed,
1201 }),
1202 };
1203 if let Some(additional_params) = request.additional_params {
1205 crate::json_utils::merge_inplace(&mut body, additional_params);
1206 }
1207
1208 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1227pub struct Models {
1228 pub provider: OpenAIConfig,
1230}
1231
1232impl Models {
1233 pub fn new(provider: OpenAIConfig) -> Self {
1235 Self { provider }
1236 }
1237}
1238
1239#[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 #[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#[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#[derive(Debug, Deserialize)]
1299pub struct ModelsReply {
1300 #[serde(default)]
1301 pub data: Vec<ModelEntry>,
1302}
1303
1304#[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 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1352pub struct Rerank {
1353 pub provider: OpenAIConfig,
1355 pub model: String,
1357 pub top_n: Option<usize>,
1360}
1361
1362impl Rerank {
1363 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 pub fn with_top_n(mut self, top_n: usize) -> Self {
1374 self.top_n = Some(top_n);
1375 self
1376 }
1377}
1378
1379#[derive(Debug, Clone, Serialize, Deserialize)]
1381pub struct RerankResultEntry {
1382 pub index: usize,
1384 #[serde(alias = "score")]
1386 pub relevance_score: f64,
1387 #[serde(default, alias = "text")]
1390 pub document: Option<String>,
1391}
1392
1393#[derive(Debug, Clone, Serialize, Deserialize, Default)]
1395pub struct RerankUsage {
1396 #[serde(default)]
1398 pub prompt_tokens: u64,
1399 #[serde(default)]
1401 pub total_tokens: u64,
1402}
1403
1404#[derive(Debug, Clone, Serialize, Deserialize)]
1406pub struct RerankReply {
1407 #[serde(default)]
1409 pub model: Option<String>,
1410 pub results: Vec<RerankResultEntry>,
1412 #[serde(default)]
1414 pub usage: Option<RerankUsage>,
1415}
1416
1417#[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 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
1517pub struct Verify {
1518 pub provider: OpenAIConfig,
1520}
1521
1522impl Verify {
1523 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;