Skip to main content

rig_core/operation/
modality.rs

1//! Buffered embedding, reranking, transcription, image, and audio operations.
2//!
3//! ```
4//! use rig_core::wire::Capabilities;
5//!
6//! let capabilities = Capabilities::embedding(32, 768).declaring(Some(768));
7//! assert_eq!(capabilities.declared, Some(768));
8//! ```
9
10use std::convert::Infallible;
11
12use super::Whole;
13use crate::embeddings::Embedding as Vector;
14use crate::error::ProviderError;
15use crate::telemetry::{GenAiOperation, SpanBuilder, SpanCombinator};
16use crate::wire::{Call, Capabilities, Fold, Free, Operation, Reply};
17
18/// A reranking request: the query, the documents to order, and the batch
19/// limit's subject.
20///
21/// The consumer trait takes `(&str, Vec<String>)`; the operation takes one
22/// request value, as every other operation does. It is also the effect bus's
23/// rerank payload.
24#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
25pub struct RerankRequest {
26    /// What the documents are ordered against.
27    pub query: String,
28    /// The documents to order, in input order.
29    pub documents: Vec<String>,
30}
31
32/// Declare one unary modality operation.
33macro_rules! modality_operation {
34    (
35        $(#[$doc:meta])*
36        $op:ident {
37            request: $request:ty,
38            response: $response:ty,
39            telemetry: $telemetry:ident,
40            fold: $fold:ty,
41            seed: $seed:expr,
42        }
43    ) => {
44        $(#[$doc])*
45        #[derive(Debug, Clone, Copy, PartialEq, Eq)]
46        pub struct $op;
47
48        impl Operation for $op {
49            type Request = $request;
50            type Event = Infallible;
51            type End = $response;
52            type Response = $response;
53            type Fold = Traced<Self, $fold>;
54            type Emit = Free;
55
56            fn fold(request: &Self::Request, call: &mut Call<'_>) -> Self::Fold {
57                let telemetry = call
58                    .wire
59                    .telemetry
60                    .map_or(GenAiOperation::$telemetry, |telemetry| telemetry(call.mode));
61                debug_assert!(!telemetry.is_completion());
62                let span = SpanBuilder::new(
63                    call.wire.name,
64                    call.wire.model.unwrap_or_default(),
65                    telemetry,
66                )
67                .build();
68                call.instrument(span.clone());
69                #[allow(clippy::redundant_closure_call)]
70                let inner = ($seed)(request, &*call);
71                Traced {
72                    inner,
73                    span,
74                    record: |span, response| {
75                        span.record_response(
76                            response.response_id.as_deref(),
77                            response.model.as_deref(),
78                            &response.usage,
79                        )
80                    },
81                }
82            }
83        }
84    };
85}
86
87/// A fold whose response is recorded on the call's telemetry span.
88pub struct Traced<Op: Operation, F> {
89    inner: F,
90    span: tracing::Span,
91    record: fn(&tracing::Span, &Op::Response),
92}
93
94impl<Op, F> Fold<Op> for Traced<Op, F>
95where
96    Op: Operation,
97    F: Fold<Op>,
98{
99    fn absorb(&mut self, event: &Op::Event) -> Result<(), ProviderError> {
100        self.inner.absorb(event)
101    }
102
103    fn finish(self, end: Op::End, reply: Reply) -> Result<Op::Response, ProviderError> {
104        let response = self.inner.finish(end, reply)?;
105        (self.record)(&self.span, &response);
106        Ok(response)
107    }
108}
109
110/// A modality response the driver writes its facts onto.
111pub(crate) trait Stamp {
112    /// Write the provider, the transport request id and the reply document,
113    /// and drop a model or id reported empty.
114    fn stamp(&mut self, reply: &Reply);
115}
116
117macro_rules! stamp {
118    ($($response:ty),*) => {$(
119        impl Stamp for $response {
120            fn stamp(&mut self, reply: &Reply) {
121                use crate::provider_response::reported;
122                self.provider.clone_from(&reply.provider);
123                self.provider_request_id = reported(reply.provider_request_id.clone());
124                self.raw.clone_from(&reply.raw);
125                self.model = reported(self.model.take());
126                self.response_id = reported(self.response_id.take());
127            }
128        }
129    )*};
130}
131
132stamp!(
133    crate::embeddings::EmbeddingResponse,
134    crate::rerank::RerankResponse,
135    crate::transcription::TranscriptionResponse
136);
137#[cfg(feature = "image")]
138stamp!(crate::image_generation::ImageGenerationResponse);
139#[cfg(feature = "audio")]
140stamp!(crate::audio_generation::AudioGenerationResponse);
141
142modality_operation!(
143    /// Embedding a batch of texts.
144    Embedding {
145        request: Vec<String>,
146        response: crate::embeddings::EmbeddingResponse,
147        telemetry: Embeddings,
148        fold: Embedded,
149        seed: |texts: &Vec<String>, call: &Call<'_>| Embedded::over(texts.clone(), call),
150    }
151);
152
153modality_operation!(
154    /// Embedding a batch of images from their encoded file bytes.
155    ImageEmbedding {
156        request: Vec<Vec<u8>>,
157        response: crate::embeddings::EmbeddingResponse,
158        telemetry: Embeddings,
159        fold: Embedded,
160        seed: |images: &Vec<Vec<u8>>, call: &Call<'_>| Embedded::over(
161            images.iter().map(|bytes| crate::embeddings::image_document(bytes)).collect(),
162            call,
163        ),
164    }
165);
166
167modality_operation!(
168    /// Ordering documents by relevance to a query.
169    Rerank {
170        request: RerankRequest,
171        response: crate::rerank::RerankResponse,
172        telemetry: Rerank,
173        fold: Whole<Self>,
174        seed: |_: &_, _: &Call<'_>| Whole::stamping(<crate::rerank::RerankResponse as Stamp>::stamp),
175    }
176);
177
178modality_operation!(
179    /// Transcribes audio using provider-specific request encoding.
180    Transcription {
181        request: crate::transcription::TranscriptionRequest,
182        response: crate::transcription::TranscriptionResponse,
183        telemetry: Transcription,
184        fold: Whole<Self>,
185        seed: |_: &_, _: &Call<'_>| Whole::stamping(<crate::transcription::TranscriptionResponse as Stamp>::stamp),
186    }
187);
188
189#[cfg(feature = "image")]
190modality_operation!(
191    /// Generating an image.
192    ImageGeneration {
193        request: crate::image_generation::ImageGenerationRequest,
194        response: crate::image_generation::ImageGenerationResponse,
195        telemetry: ImageGeneration,
196        fold: Whole<Self>,
197        seed: |_: &_, _: &Call<'_>| Whole::stamping(<crate::image_generation::ImageGenerationResponse as Stamp>::stamp),
198    }
199);
200
201#[cfg(feature = "audio")]
202modality_operation!(
203    /// Generating speech.
204    AudioGeneration {
205        request: crate::audio_generation::AudioGenerationRequest,
206        response: crate::audio_generation::AudioGenerationResponse,
207        telemetry: AudioGeneration,
208        fold: Whole<Self>,
209        seed: |_: &_, _: &Call<'_>| Whole::stamping(<crate::audio_generation::AudioGenerationResponse as Stamp>::stamp),
210    }
211);
212
213/// Pairs a reply's vectors positionally with the request's documents.
214/// Finishing rejects unequal vector and document counts, and widths that
215/// contradict the width the caller declared.
216pub struct Embedded {
217    documents: Vec<String>,
218    /// The provider and capabilities the reply is checked against.
219    provider: String,
220    capabilities: Capabilities,
221}
222
223impl Embedded {
224    /// A fold that will zip its vectors onto `documents`, for the model
225    /// `call` describes.
226    pub fn over(documents: Vec<String>, call: &Call<'_>) -> Self {
227        Self {
228            documents,
229            provider: call.wire.name.to_owned(),
230            capabilities: call.wire.capabilities,
231        }
232    }
233
234    /// The reply's vectors, paired with the inputs they belong to.
235    fn zipped(self, vectors: Vec<Vector>) -> Result<Vec<Vector>, ProviderError> {
236        self.capabilities.honour_declaration(
237            &self.provider,
238            vectors.iter().map(|vector| vector.vec.len()),
239        )?;
240        if vectors.len() != self.documents.len() {
241            return Err(ProviderError::Response(format!(
242                "provider returned {} embeddings for {} documents",
243                vectors.len(),
244                self.documents.len()
245            )));
246        }
247        Ok(self
248            .documents
249            .into_iter()
250            .zip(vectors)
251            .map(|(document, vector)| Vector {
252                document,
253                vec: vector.vec,
254            })
255            .collect())
256    }
257}
258
259impl<Op> Fold<Op> for Embedded
260where
261    Op: Operation<
262            Event = Infallible,
263            End = crate::embeddings::EmbeddingResponse,
264            Response = crate::embeddings::EmbeddingResponse,
265        >,
266{
267    fn absorb(&mut self, event: &Infallible) -> Result<(), ProviderError> {
268        match *event {}
269    }
270
271    fn finish(
272        self,
273        mut response: crate::embeddings::EmbeddingResponse,
274        reply: Reply,
275    ) -> Result<crate::embeddings::EmbeddingResponse, ProviderError> {
276        response.embeddings = self.zipped(std::mem::take(&mut response.embeddings))?;
277        response.stamp(&reply);
278        Ok(response)
279    }
280}