1use 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#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
25pub struct RerankRequest {
26 pub query: String,
28 pub documents: Vec<String>,
30}
31
32macro_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
87pub 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
110pub(crate) trait Stamp {
112 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 {
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 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 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 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 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 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
213pub struct Embedded {
217 documents: Vec<String>,
218 provider: String,
220 capabilities: Capabilities,
221}
222
223impl Embedded {
224 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 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}