ferrin_core/
modality_hooks.rs1use std::fmt;
7
8use ferrin_spec::Headers;
9use ferrin_spec::JsonValue;
10use ferrin_spec::ProviderMetadata;
11use ferrin_spec::ProviderOptions;
12use ferrin_spec::ResponseMetadata;
13use ferrin_spec::Warning;
14use serde_json::json;
15
16use crate::embed::Embedding;
17use crate::embed::EmbeddingUsage;
18use crate::hooks::HookList;
19use crate::rerank::Ranked;
20use crate::rerank::RerankDocument;
21use crate::telemetry::ModelIdentity;
22
23#[derive(Debug, Clone, PartialEq)]
25#[non_exhaustive]
26pub enum EmbeddingInput {
27 Single(String),
29 Many(Vec<String>),
31}
32
33#[derive(Debug, Clone, PartialEq)]
35#[non_exhaustive]
36pub enum EmbeddingOutput {
37 Single(Embedding),
39 Many(Vec<Embedding>),
41}
42
43#[derive(Debug, Clone, PartialEq)]
45#[non_exhaustive]
46pub enum EmbeddingResponse {
47 Single(Box<ResponseMetadata>),
49 Many(Vec<ResponseMetadata>),
51}
52
53#[derive(Debug, Clone, PartialEq)]
55pub struct EmbedCallStartEvent {
56 pub runtime_context: Option<JsonValue>,
58 pub call_id: String,
60 pub operation_id: &'static str,
62 pub model: ModelIdentity,
64 pub value: Option<EmbeddingInput>,
66 pub max_retries: u32,
68 pub headers: Headers,
70 pub provider_options: ProviderOptions,
72}
73
74#[derive(Debug, Clone, PartialEq)]
76pub struct EmbedCallEndEvent {
77 pub runtime_context: Option<JsonValue>,
79 pub call_id: String,
81 pub operation_id: &'static str,
83 pub model: ModelIdentity,
85 pub value: Option<EmbeddingInput>,
87 pub embedding: Option<EmbeddingOutput>,
89 pub usage: EmbeddingUsage,
91 pub warnings: Vec<Warning>,
93 pub provider_metadata: Option<ProviderMetadata>,
95 pub response: EmbeddingResponse,
97}
98
99#[derive(Debug, Clone, PartialEq)]
101pub struct RerankCallStartEvent {
102 pub runtime_context: Option<JsonValue>,
104 pub call_id: String,
106 pub operation_id: &'static str,
108 pub model: ModelIdentity,
110 pub documents: Option<Vec<RerankDocument>>,
112 pub query: Option<String>,
114 pub top_n: Option<usize>,
116 pub max_retries: u32,
118 pub headers: Headers,
120 pub provider_options: ProviderOptions,
122}
123
124#[derive(Debug, Clone, PartialEq)]
126pub struct RerankCallEndEvent {
127 pub runtime_context: Option<JsonValue>,
129 pub call_id: String,
131 pub operation_id: &'static str,
133 pub model: ModelIdentity,
135 pub documents: Option<Vec<RerankDocument>>,
137 pub query: Option<String>,
139 pub ranking: Option<Vec<Ranked<RerankDocument>>>,
141 pub warnings: Vec<Warning>,
143 pub provider_metadata: Option<ProviderMetadata>,
145 pub response: ResponseMetadata,
147}
148
149pub(crate) struct ModalityHooks<S, E> {
150 pub(crate) runtime_context: JsonValue,
151 pub(crate) on_start: HookList<S>,
152 pub(crate) on_end: HookList<E>,
153}
154
155impl<S, E> Default for ModalityHooks<S, E> {
156 fn default() -> Self {
157 Self {
158 runtime_context: json!({}),
159 on_start: Vec::new(),
160 on_end: Vec::new(),
161 }
162 }
163}
164
165impl<S, E> fmt::Debug for ModalityHooks<S, E> {
166 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
167 f.debug_struct("ModalityHooks")
168 .field("on_start", &self.on_start.len())
169 .field("on_end", &self.on_end.len())
170 .finish_non_exhaustive()
171 }
172}
173
174macro_rules! impl_modality_hooks {
175 ($ty:ident $(<$generic:ident>)?, $start:ty, $end:ty) => {
176 impl$(<$generic>)? $ty$(<$generic>)? {
177 #[must_use]
179 pub fn runtime_context(mut self, context: ::ferrin_spec::JsonValue) -> Self {
180 self.hooks.runtime_context = context;
181 self
182 }
183
184 #[must_use]
186 pub fn on_start(mut self, hook: impl $crate::hooks::HookFn<$start>) -> Self {
187 self.hooks.on_start.push(::std::sync::Arc::new(hook));
188 self
189 }
190
191 #[must_use]
193 pub fn on_end(mut self, hook: impl $crate::hooks::HookFn<$end>) -> Self {
194 self.hooks.on_end.push(::std::sync::Arc::new(hook));
195 self
196 }
197 }
198 };
199}
200pub(crate) use impl_modality_hooks;