Skip to main content

starweaver_model/wrappers/
hooked.rs

1//! Model execution hook wrapper.
2
3use std::sync::Arc;
4
5use async_trait::async_trait;
6use serde::{Deserialize, Serialize};
7use serde_json::{Map, Value};
8use starweaver_core::{ConversationId, RunId};
9use starweaver_usage::Usage;
10
11use super::DynModelAdapter;
12use crate::{
13    adapter::{
14        ModelAdapter, ModelError, ModelRequestContext, ModelRequestParameters,
15        ModelResponseEventStream, ModelRunSession,
16    },
17    message::{ModelMessage, ModelResponse},
18    profile::ModelProfile,
19    settings::ModelSettings,
20    stream::ModelResponseStreamEvent,
21};
22
23/// Shared model execution hook.
24pub type DynModelExecutionHook = Arc<dyn ModelExecutionHook>;
25
26/// Model wrapper metadata passed to execution hooks.
27#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
28pub struct ModelExecutionMetadata {
29    /// Provider model name.
30    pub model_name: String,
31    /// Provider name when available.
32    #[serde(default, skip_serializing_if = "Option::is_none")]
33    pub provider_name: Option<String>,
34    /// Run id.
35    pub run_id: RunId,
36    /// Conversation id.
37    pub conversation_id: ConversationId,
38    /// Agent id when present in trace metadata.
39    #[serde(default, skip_serializing_if = "Option::is_none")]
40    pub agent_id: Option<String>,
41    /// Agent name when present in trace metadata.
42    #[serde(default, skip_serializing_if = "Option::is_none")]
43    pub agent_name: Option<String>,
44    /// Whether the request uses the streaming path.
45    pub stream: bool,
46    /// Low-cardinality context metadata copied from the runtime request context.
47    #[serde(default, skip_serializing_if = "Map::is_empty")]
48    pub context_metadata: Map<String, Value>,
49}
50
51impl ModelExecutionMetadata {
52    fn new(model: &dyn ModelAdapter, context: &ModelRequestContext, stream: bool) -> Self {
53        let agent_id = context
54            .llm_trace_metadata
55            .get("agent_id")
56            .or_else(|| context.llm_trace_metadata.get("starweaver.agent_id"))
57            .and_then(Value::as_str)
58            .map(ToString::to_string);
59        let agent_name = context
60            .llm_trace_metadata
61            .get("agent_name")
62            .or_else(|| context.llm_trace_metadata.get("starweaver.agent_name"))
63            .and_then(Value::as_str)
64            .map(ToString::to_string);
65        Self {
66            model_name: model.model_name().to_string(),
67            provider_name: model.provider_name().map(ToString::to_string),
68            run_id: context.run_id.clone(),
69            conversation_id: context.conversation_id.clone(),
70            agent_id,
71            agent_name,
72            stream,
73            context_metadata: context.llm_trace_metadata.clone(),
74        }
75    }
76}
77
78/// Hook around model adapter execution.
79#[async_trait]
80pub trait ModelExecutionHook: Send + Sync {
81    /// Called before the wrapped model adapter receives a request.
82    ///
83    /// # Errors
84    ///
85    /// Returning an error prevents the inner model request from executing.
86    async fn before_model_request(
87        &self,
88        _metadata: ModelExecutionMetadata,
89        _messages: &[ModelMessage],
90        _settings: Option<&ModelSettings>,
91        _params: &ModelRequestParameters,
92        _context: &ModelRequestContext,
93    ) -> Result<(), ModelError> {
94        Ok(())
95    }
96
97    /// Called after the wrapped model adapter returns a final response.
98    ///
99    /// # Errors
100    ///
101    /// Returning an error fails the model request.
102    async fn after_model_response(
103        &self,
104        _metadata: ModelExecutionMetadata,
105        _response: &ModelResponse,
106    ) -> Result<(), ModelError> {
107        Ok(())
108    }
109
110    /// Called when the wrapped model adapter returns an error.
111    ///
112    /// # Errors
113    ///
114    /// Returning an error replaces the original model error.
115    async fn on_model_error(
116        &self,
117        _metadata: ModelExecutionMetadata,
118        _error: &ModelError,
119    ) -> Result<(), ModelError> {
120        Ok(())
121    }
122}
123
124/// Model wrapper that exposes request/response lifecycle hooks.
125pub struct HookedModel {
126    inner: DynModelAdapter,
127    hooks: Vec<DynModelExecutionHook>,
128}
129
130struct HookedModelRunSession<'a> {
131    model: &'a HookedModel,
132    inner: Box<dyn ModelRunSession + 'a>,
133}
134
135impl HookedModel {
136    /// Create a hooked model wrapper.
137    #[must_use]
138    pub fn new(inner: DynModelAdapter) -> Self {
139        Self {
140            inner,
141            hooks: Vec::new(),
142        }
143    }
144
145    /// Add one execution hook.
146    #[must_use]
147    pub fn with_hook(mut self, hook: DynModelExecutionHook) -> Self {
148        self.hooks.push(hook);
149        self
150    }
151
152    async fn call_before(
153        &self,
154        metadata: &ModelExecutionMetadata,
155        messages: &[ModelMessage],
156        settings: Option<&ModelSettings>,
157        params: &ModelRequestParameters,
158        context: &ModelRequestContext,
159    ) -> Result<(), ModelError> {
160        for hook in &self.hooks {
161            hook.before_model_request(metadata.clone(), messages, settings, params, context)
162                .await?;
163        }
164        Ok(())
165    }
166
167    async fn call_after(
168        hooks: &[DynModelExecutionHook],
169        metadata: &ModelExecutionMetadata,
170        response: &ModelResponse,
171    ) -> Result<(), ModelError> {
172        for hook in hooks {
173            hook.after_model_response(metadata.clone(), response)
174                .await?;
175        }
176        Ok(())
177    }
178
179    async fn call_error(
180        hooks: &[DynModelExecutionHook],
181        metadata: &ModelExecutionMetadata,
182        error: &ModelError,
183    ) -> Result<(), ModelError> {
184        for hook in hooks {
185            hook.on_model_error(metadata.clone(), error).await?;
186        }
187        Ok(())
188    }
189}
190
191#[async_trait]
192impl ModelAdapter for HookedModel {
193    fn model_name(&self) -> &str {
194        self.inner.model_name()
195    }
196
197    fn provider_name(&self) -> Option<&str> {
198        self.inner.provider_name()
199    }
200
201    fn profile(&self) -> &ModelProfile {
202        self.inner.profile()
203    }
204
205    fn default_settings(&self) -> Option<&ModelSettings> {
206        self.inner.default_settings()
207    }
208
209    fn start_run_session(&self) -> Box<dyn ModelRunSession + '_> {
210        Box::new(HookedModelRunSession {
211            model: self,
212            inner: self.inner.start_run_session(),
213        })
214    }
215
216    async fn request(
217        &self,
218        messages: Vec<ModelMessage>,
219        settings: Option<ModelSettings>,
220        params: ModelRequestParameters,
221        context: ModelRequestContext,
222    ) -> Result<ModelResponse, ModelError> {
223        let metadata = ModelExecutionMetadata::new(self.inner.as_ref(), &context, false);
224        self.call_before(&metadata, &messages, settings.as_ref(), &params, &context)
225            .await?;
226        match self
227            .inner
228            .request(messages, settings, params, context)
229            .await
230        {
231            Ok(response) => {
232                Self::call_after(&self.hooks, &metadata, &response).await?;
233                Ok(response)
234            }
235            Err(error) => {
236                Self::call_error(&self.hooks, &metadata, &error).await?;
237                Err(error)
238            }
239        }
240    }
241
242    async fn request_stream(
243        &self,
244        messages: Vec<ModelMessage>,
245        settings: Option<ModelSettings>,
246        params: ModelRequestParameters,
247        context: ModelRequestContext,
248    ) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
249        let metadata = ModelExecutionMetadata::new(self.inner.as_ref(), &context, true);
250        self.call_before(&metadata, &messages, settings.as_ref(), &params, &context)
251            .await?;
252        match self
253            .inner
254            .request_stream(messages, settings, params, context)
255            .await
256        {
257            Ok(events) => {
258                if let Some(response) = events.iter().find_map(|event| match event {
259                    ModelResponseStreamEvent::FinalResult(response) => Some(response.as_ref()),
260                    ModelResponseStreamEvent::PartStart(_)
261                    | ModelResponseStreamEvent::PartDelta(_)
262                    | ModelResponseStreamEvent::PartEnd(_)
263                    | ModelResponseStreamEvent::Diagnostic(_) => None,
264                }) {
265                    Self::call_after(&self.hooks, &metadata, response).await?;
266                }
267                Ok(events)
268            }
269            Err(error) => {
270                Self::call_error(&self.hooks, &metadata, &error).await?;
271                Err(error)
272            }
273        }
274    }
275
276    async fn request_stream_incremental(
277        &self,
278        messages: Vec<ModelMessage>,
279        settings: Option<ModelSettings>,
280        params: ModelRequestParameters,
281        context: ModelRequestContext,
282    ) -> Result<ModelResponseEventStream, ModelError> {
283        let metadata = ModelExecutionMetadata::new(self.inner.as_ref(), &context, true);
284        self.call_before(&metadata, &messages, settings.as_ref(), &params, &context)
285            .await?;
286        match self
287            .inner
288            .request_stream_incremental(messages, settings, params, context)
289            .await
290        {
291            Ok(mut inner_stream) => {
292                let drop_abort_token = inner_stream.drop_abort_token();
293                let hooks = self.hooks.clone();
294                let (sender, receiver) = tokio::sync::mpsc::channel(32);
295                tokio::spawn(async move {
296                    while let Some(event) = inner_stream.recv().await {
297                        match event {
298                            Ok(ModelResponseStreamEvent::FinalResult(response)) => {
299                                if let Err(error) =
300                                    Self::call_after(&hooks, &metadata, &response).await
301                                {
302                                    let _ = sender.send(Err(error)).await;
303                                    return;
304                                }
305                                if sender
306                                    .send(Ok(ModelResponseStreamEvent::FinalResult(response)))
307                                    .await
308                                    .is_err()
309                                {
310                                    return;
311                                }
312                            }
313                            Ok(event) => {
314                                if sender.send(Ok(event)).await.is_err() {
315                                    return;
316                                }
317                            }
318                            Err(error) => {
319                                let replacement =
320                                    Self::call_error(&hooks, &metadata, &error).await.err();
321                                let _ = sender.send(Err(replacement.unwrap_or(error))).await;
322                                return;
323                            }
324                        }
325                    }
326                });
327                Ok(
328                    ModelResponseEventStream::new_with_cancellation_and_drop_abort(
329                        receiver,
330                        starweaver_core::CancellationToken::default(),
331                        drop_abort_token,
332                    ),
333                )
334            }
335            Err(error) => {
336                Self::call_error(&self.hooks, &metadata, &error).await?;
337                Err(error)
338            }
339        }
340    }
341
342    async fn count_tokens(
343        &self,
344        messages: &[ModelMessage],
345        settings: Option<&ModelSettings>,
346        params: &ModelRequestParameters,
347    ) -> Result<Usage, ModelError> {
348        self.inner.count_tokens(messages, settings, params).await
349    }
350}
351
352#[async_trait]
353impl ModelRunSession for HookedModelRunSession<'_> {
354    async fn request_stream_incremental(
355        &mut self,
356        messages: Vec<ModelMessage>,
357        settings: Option<ModelSettings>,
358        params: ModelRequestParameters,
359        context: ModelRequestContext,
360    ) -> Result<ModelResponseEventStream, ModelError> {
361        let cancellation_token = context.cancellation_token();
362        let metadata = ModelExecutionMetadata::new(self.model.inner.as_ref(), &context, true);
363        self.model
364            .call_before(&metadata, &messages, settings.as_ref(), &params, &context)
365            .await?;
366        match self
367            .inner
368            .request_stream_incremental(messages, settings, params, context)
369            .await
370        {
371            Ok(mut inner_stream) => {
372                let drop_abort_token = inner_stream.drop_abort_token();
373                let hooks = self.model.hooks.clone();
374                let (sender, receiver) = tokio::sync::mpsc::channel(32);
375                tokio::spawn(async move {
376                    while let Some(event) = inner_stream.recv().await {
377                        match event {
378                            Ok(ModelResponseStreamEvent::FinalResult(response)) => {
379                                if let Err(error) =
380                                    HookedModel::call_after(&hooks, &metadata, &response).await
381                                {
382                                    let _ = sender.send(Err(error)).await;
383                                    return;
384                                }
385                                if sender
386                                    .send(Ok(ModelResponseStreamEvent::FinalResult(response)))
387                                    .await
388                                    .is_err()
389                                {
390                                    return;
391                                }
392                            }
393                            Ok(event) => {
394                                if sender.send(Ok(event)).await.is_err() {
395                                    return;
396                                }
397                            }
398                            Err(error) => {
399                                let replacement =
400                                    HookedModel::call_error(&hooks, &metadata, &error)
401                                        .await
402                                        .err();
403                                let _ = sender.send(Err(replacement.unwrap_or(error))).await;
404                                return;
405                            }
406                        }
407                    }
408                });
409                Ok(
410                    ModelResponseEventStream::new_with_cancellation_and_drop_abort(
411                        receiver,
412                        cancellation_token,
413                        drop_abort_token,
414                    ),
415                )
416            }
417            Err(error) => {
418                HookedModel::call_error(&self.model.hooks, &metadata, &error).await?;
419                Err(error)
420            }
421        }
422    }
423
424    async fn close(&mut self) {
425        self.inner.close().await;
426    }
427}