use std::sync::Arc;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use starweaver_core::{ConversationId, RunId};
use starweaver_usage::Usage;
use super::DynModelAdapter;
use crate::{
adapter::{
ModelAdapter, ModelError, ModelRequestContext, ModelRequestParameters,
ModelResponseEventStream,
},
message::{ModelMessage, ModelResponse},
profile::ModelProfile,
settings::ModelSettings,
stream::ModelResponseStreamEvent,
};
pub type DynModelExecutionHook = Arc<dyn ModelExecutionHook>;
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ModelExecutionMetadata {
pub model_name: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_name: Option<String>,
pub run_id: RunId,
pub conversation_id: ConversationId,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub agent_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub agent_name: Option<String>,
pub stream: bool,
#[serde(default, skip_serializing_if = "Map::is_empty")]
pub context_metadata: Map<String, Value>,
}
impl ModelExecutionMetadata {
fn new(model: &dyn ModelAdapter, context: &ModelRequestContext, stream: bool) -> Self {
let agent_id = context
.llm_trace_metadata
.get("agent_id")
.or_else(|| context.llm_trace_metadata.get("starweaver.agent_id"))
.and_then(Value::as_str)
.map(ToString::to_string);
let agent_name = context
.llm_trace_metadata
.get("agent_name")
.or_else(|| context.llm_trace_metadata.get("starweaver.agent_name"))
.and_then(Value::as_str)
.map(ToString::to_string);
Self {
model_name: model.model_name().to_string(),
provider_name: model.provider_name().map(ToString::to_string),
run_id: context.run_id.clone(),
conversation_id: context.conversation_id.clone(),
agent_id,
agent_name,
stream,
context_metadata: context.llm_trace_metadata.clone(),
}
}
}
#[async_trait]
pub trait ModelExecutionHook: Send + Sync {
async fn before_model_request(
&self,
_metadata: ModelExecutionMetadata,
_messages: &[ModelMessage],
_settings: Option<&ModelSettings>,
_params: &ModelRequestParameters,
_context: &ModelRequestContext,
) -> Result<(), ModelError> {
Ok(())
}
async fn after_model_response(
&self,
_metadata: ModelExecutionMetadata,
_response: &ModelResponse,
) -> Result<(), ModelError> {
Ok(())
}
async fn on_model_error(
&self,
_metadata: ModelExecutionMetadata,
_error: &ModelError,
) -> Result<(), ModelError> {
Ok(())
}
}
pub struct HookedModel {
inner: DynModelAdapter,
hooks: Vec<DynModelExecutionHook>,
}
impl HookedModel {
#[must_use]
pub fn new(inner: DynModelAdapter) -> Self {
Self {
inner,
hooks: Vec::new(),
}
}
#[must_use]
pub fn with_hook(mut self, hook: DynModelExecutionHook) -> Self {
self.hooks.push(hook);
self
}
async fn call_before(
&self,
metadata: &ModelExecutionMetadata,
messages: &[ModelMessage],
settings: Option<&ModelSettings>,
params: &ModelRequestParameters,
context: &ModelRequestContext,
) -> Result<(), ModelError> {
for hook in &self.hooks {
hook.before_model_request(metadata.clone(), messages, settings, params, context)
.await?;
}
Ok(())
}
async fn call_after(
hooks: &[DynModelExecutionHook],
metadata: &ModelExecutionMetadata,
response: &ModelResponse,
) -> Result<(), ModelError> {
for hook in hooks {
hook.after_model_response(metadata.clone(), response)
.await?;
}
Ok(())
}
async fn call_error(
hooks: &[DynModelExecutionHook],
metadata: &ModelExecutionMetadata,
error: &ModelError,
) -> Result<(), ModelError> {
for hook in hooks {
hook.on_model_error(metadata.clone(), error).await?;
}
Ok(())
}
}
#[async_trait]
impl ModelAdapter for HookedModel {
fn model_name(&self) -> &str {
self.inner.model_name()
}
fn provider_name(&self) -> Option<&str> {
self.inner.provider_name()
}
fn profile(&self) -> &ModelProfile {
self.inner.profile()
}
fn default_settings(&self) -> Option<&ModelSettings> {
self.inner.default_settings()
}
async fn request(
&self,
messages: Vec<ModelMessage>,
settings: Option<ModelSettings>,
params: ModelRequestParameters,
context: ModelRequestContext,
) -> Result<ModelResponse, ModelError> {
let metadata = ModelExecutionMetadata::new(self.inner.as_ref(), &context, false);
self.call_before(&metadata, &messages, settings.as_ref(), ¶ms, &context)
.await?;
match self
.inner
.request(messages, settings, params, context)
.await
{
Ok(response) => {
Self::call_after(&self.hooks, &metadata, &response).await?;
Ok(response)
}
Err(error) => {
Self::call_error(&self.hooks, &metadata, &error).await?;
Err(error)
}
}
}
async fn request_stream(
&self,
messages: Vec<ModelMessage>,
settings: Option<ModelSettings>,
params: ModelRequestParameters,
context: ModelRequestContext,
) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
let metadata = ModelExecutionMetadata::new(self.inner.as_ref(), &context, true);
self.call_before(&metadata, &messages, settings.as_ref(), ¶ms, &context)
.await?;
match self
.inner
.request_stream(messages, settings, params, context)
.await
{
Ok(events) => {
if let Some(response) = events.iter().find_map(|event| match event {
ModelResponseStreamEvent::FinalResult(response) => Some(response.as_ref()),
ModelResponseStreamEvent::PartStart(_)
| ModelResponseStreamEvent::PartDelta(_)
| ModelResponseStreamEvent::PartEnd(_) => None,
}) {
Self::call_after(&self.hooks, &metadata, response).await?;
}
Ok(events)
}
Err(error) => {
Self::call_error(&self.hooks, &metadata, &error).await?;
Err(error)
}
}
}
async fn request_stream_incremental(
&self,
messages: Vec<ModelMessage>,
settings: Option<ModelSettings>,
params: ModelRequestParameters,
context: ModelRequestContext,
) -> Result<ModelResponseEventStream, ModelError> {
let metadata = ModelExecutionMetadata::new(self.inner.as_ref(), &context, true);
self.call_before(&metadata, &messages, settings.as_ref(), ¶ms, &context)
.await?;
match self
.inner
.request_stream_incremental(messages, settings, params, context)
.await
{
Ok(mut inner_stream) => {
let hooks = self.hooks.clone();
let (sender, receiver) = tokio::sync::mpsc::channel(32);
tokio::spawn(async move {
while let Some(event) = inner_stream.recv().await {
match event {
Ok(ModelResponseStreamEvent::FinalResult(response)) => {
if let Err(error) =
Self::call_after(&hooks, &metadata, &response).await
{
let _ = sender.send(Err(error)).await;
return;
}
if sender
.send(Ok(ModelResponseStreamEvent::FinalResult(response)))
.await
.is_err()
{
return;
}
}
Ok(event) => {
if sender.send(Ok(event)).await.is_err() {
return;
}
}
Err(error) => {
let replacement =
Self::call_error(&hooks, &metadata, &error).await.err();
let _ = sender.send(Err(replacement.unwrap_or(error))).await;
return;
}
}
}
});
Ok(ModelResponseEventStream::new(receiver))
}
Err(error) => {
Self::call_error(&self.hooks, &metadata, &error).await?;
Err(error)
}
}
}
async fn count_tokens(
&self,
messages: &[ModelMessage],
settings: Option<&ModelSettings>,
params: &ModelRequestParameters,
) -> Result<Usage, ModelError> {
self.inner.count_tokens(messages, settings, params).await
}
}