use std::marker::PhantomData;
use std::sync::Arc;
use rig::agent::{HookAction, PromptHook, ToolCallHookAction};
use rig::completion::{CompletionModel, CompletionResponse, Message};
use crate::emit::emit_kind;
use crate::event::{EventKind, PAYLOAD_TRUNCATE_BYTES, truncate_utf8};
pub type ConversationIdResolver = Arc<dyn Fn() -> Option<String> + Send + Sync>;
pub type ModelResolver<R> = Arc<dyn Fn(&CompletionResponse<R>) -> Option<String> + Send + Sync>;
#[derive(Debug, Clone)]
pub struct TelemetryHookConfig {
pub model: String,
pub conversation_id: String,
pub payload_truncate_bytes: usize,
}
impl TelemetryHookConfig {
pub fn new(model: impl Into<String>, conversation_id: impl Into<String>) -> Self {
Self {
model: model.into(),
conversation_id: conversation_id.into(),
payload_truncate_bytes: PAYLOAD_TRUNCATE_BYTES,
}
}
}
pub struct TelemetryHook<M: CompletionModel> {
config: TelemetryHookConfig,
conversation_id_resolver: Option<ConversationIdResolver>,
model_resolver: Option<ModelResolver<M::Response>>,
_model: PhantomData<fn() -> M>,
}
impl<M: CompletionModel> TelemetryHook<M> {
pub fn new(config: TelemetryHookConfig) -> Self {
Self {
config,
conversation_id_resolver: None,
model_resolver: None,
_model: PhantomData,
}
}
pub fn with_defaults(model: impl Into<String>, conversation_id: impl Into<String>) -> Self {
Self::new(TelemetryHookConfig::new(model, conversation_id))
}
#[must_use]
pub fn with_conversation_id_resolver<F>(mut self, resolver: F) -> Self
where
F: Fn() -> Option<String> + Send + Sync + 'static,
{
self.conversation_id_resolver = Some(Arc::new(resolver));
self
}
#[must_use]
pub fn with_model_resolver<F>(mut self, resolver: F) -> Self
where
F: Fn(&CompletionResponse<M::Response>) -> Option<String> + Send + Sync + 'static,
{
self.model_resolver = Some(Arc::new(resolver));
self
}
fn resolved_conversation_id(&self) -> String {
self.conversation_id_resolver
.as_ref()
.and_then(|f| f())
.unwrap_or_else(|| self.config.conversation_id.clone())
}
fn resolved_model(&self, response: &CompletionResponse<M::Response>) -> String {
self.model_resolver
.as_ref()
.and_then(|f| f(response))
.unwrap_or_else(|| self.config.model.clone())
}
}
impl<M: CompletionModel> Clone for TelemetryHook<M> {
fn clone(&self) -> Self {
Self {
config: self.config.clone(),
conversation_id_resolver: self.conversation_id_resolver.clone(),
model_resolver: self.model_resolver.clone(),
_model: PhantomData,
}
}
}
impl<M: CompletionModel> std::fmt::Debug for TelemetryHook<M> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TelemetryHook")
.field("config", &self.config)
.field(
"conversation_id_resolver",
&self.conversation_id_resolver.as_ref().map(|_| "<fn>"),
)
.field(
"model_resolver",
&self.model_resolver.as_ref().map(|_| "<fn>"),
)
.finish_non_exhaustive()
}
}
impl<M> PromptHook<M> for TelemetryHook<M>
where
M: CompletionModel,
{
async fn on_completion_call(&self, _prompt: &Message, history: &[Message]) -> HookAction {
let messages_in = history.len().saturating_add(1);
emit_kind(
self.resolved_conversation_id(),
EventKind::PromptStarted {
model: self.config.model.clone(),
messages_in,
},
);
HookAction::cont()
}
async fn on_completion_response(
&self,
_prompt: &Message,
response: &CompletionResponse<M::Response>,
) -> HookAction {
let usage = response.usage;
emit_kind(
self.resolved_conversation_id(),
EventKind::PromptCompleted {
model: self.resolved_model(response),
tokens_in: positive(usage.input_tokens),
tokens_out: positive(usage.output_tokens),
response_id: response.message_id.clone(),
},
);
HookAction::cont()
}
async fn on_tool_call(
&self,
tool_name: &str,
tool_call_id: Option<String>,
internal_call_id: &str,
args: &str,
) -> ToolCallHookAction {
let (args_json, truncated) = truncate_utf8(args, self.config.payload_truncate_bytes);
emit_kind(
self.resolved_conversation_id(),
EventKind::ToolInvoked {
tool_name: tool_name.to_string(),
provider_call_id: tool_call_id,
call_id: internal_call_id.to_string(),
args_json,
truncated,
},
);
ToolCallHookAction::cont()
}
async fn on_tool_result(
&self,
tool_name: &str,
tool_call_id: Option<String>,
internal_call_id: &str,
_args: &str,
result: &str,
) -> HookAction {
let (result, truncated) = truncate_utf8(result, self.config.payload_truncate_bytes);
emit_kind(
self.resolved_conversation_id(),
EventKind::ToolCompleted {
tool_name: tool_name.to_string(),
provider_call_id: tool_call_id,
call_id: internal_call_id.to_string(),
result,
truncated,
},
);
HookAction::cont()
}
}
fn positive(value: u64) -> Option<u64> {
if value == 0 { None } else { Some(value) }
}