1use 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
23pub type DynModelExecutionHook = Arc<dyn ModelExecutionHook>;
25
26#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
28pub struct ModelExecutionMetadata {
29 pub model_name: String,
31 #[serde(default, skip_serializing_if = "Option::is_none")]
33 pub provider_name: Option<String>,
34 pub run_id: RunId,
36 pub conversation_id: ConversationId,
38 #[serde(default, skip_serializing_if = "Option::is_none")]
40 pub agent_id: Option<String>,
41 #[serde(default, skip_serializing_if = "Option::is_none")]
43 pub agent_name: Option<String>,
44 pub stream: bool,
46 #[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#[async_trait]
80pub trait ModelExecutionHook: Send + Sync {
81 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 async fn after_model_response(
103 &self,
104 _metadata: ModelExecutionMetadata,
105 _response: &ModelResponse,
106 ) -> Result<(), ModelError> {
107 Ok(())
108 }
109
110 async fn on_model_error(
116 &self,
117 _metadata: ModelExecutionMetadata,
118 _error: &ModelError,
119 ) -> Result<(), ModelError> {
120 Ok(())
121 }
122}
123
124pub 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 #[must_use]
138 pub fn new(inner: DynModelAdapter) -> Self {
139 Self {
140 inner,
141 hooks: Vec::new(),
142 }
143 }
144
145 #[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(), ¶ms, &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(), ¶ms, &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(), ¶ms, &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(), ¶ms, &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}