use std::sync::Arc;
use neuron_types::{
ContentBlock, HookAction, HookError, HookEvent, ObservabilityHook, WasmCompatSend,
};
use crate::guardrail::{ErasedInputGuardrail, ErasedOutputGuardrail, GuardrailResult};
pub struct GuardrailHook {
input_guardrails: Vec<Arc<dyn ErasedInputGuardrail>>,
output_guardrails: Vec<Arc<dyn ErasedOutputGuardrail>>,
}
impl GuardrailHook {
#[must_use]
pub fn new() -> Self {
Self {
input_guardrails: Vec::new(),
output_guardrails: Vec::new(),
}
}
#[must_use]
pub fn input_guardrail<G>(mut self, guardrail: G) -> Self
where
G: ErasedInputGuardrail + 'static,
{
self.input_guardrails.push(Arc::new(guardrail));
self
}
#[must_use]
pub fn output_guardrail<G>(mut self, guardrail: G) -> Self
where
G: ErasedOutputGuardrail + 'static,
{
self.output_guardrails.push(Arc::new(guardrail));
self
}
}
impl Default for GuardrailHook {
fn default() -> Self {
Self::new()
}
}
fn extract_last_user_text(messages: &[neuron_types::Message]) -> String {
for message in messages.iter().rev() {
if message.role == neuron_types::Role::User {
let texts: Vec<&str> = message
.content
.iter()
.filter_map(|block| match block {
ContentBlock::Text(t) => Some(t.as_str()),
_ => None,
})
.collect();
if !texts.is_empty() {
return texts.join("\n");
}
}
}
String::new()
}
fn extract_response_text(message: &neuron_types::Message) -> String {
let texts: Vec<&str> = message
.content
.iter()
.filter_map(|block| match block {
ContentBlock::Text(t) => Some(t.as_str()),
_ => None,
})
.collect();
texts.join("\n")
}
fn map_guardrail_result(result: GuardrailResult, direction: &str) -> HookAction {
match result {
GuardrailResult::Pass => HookAction::Continue,
GuardrailResult::Tripwire(reason) => HookAction::Terminate { reason },
GuardrailResult::Warn(reason) => {
tracing::warn!("{direction} guardrail warning: {reason}");
HookAction::Continue
}
}
}
impl ObservabilityHook for GuardrailHook {
fn on_event(
&self,
event: HookEvent<'_>,
) -> impl Future<Output = Result<HookAction, HookError>> + WasmCompatSend {
let input_guardrails = &self.input_guardrails;
let output_guardrails = &self.output_guardrails;
async move {
match event {
HookEvent::PreLlmCall { request } => {
if input_guardrails.is_empty() {
return Ok(HookAction::Continue);
}
let text = extract_last_user_text(&request.messages);
if text.is_empty() {
return Ok(HookAction::Continue);
}
for guardrail in input_guardrails {
let result = guardrail.check_dyn(&text).await;
if !result.is_pass() {
return Ok(map_guardrail_result(result, "input"));
}
}
Ok(HookAction::Continue)
}
HookEvent::PostLlmCall { response } => {
if output_guardrails.is_empty() {
return Ok(HookAction::Continue);
}
let text = extract_response_text(&response.message);
if text.is_empty() {
return Ok(HookAction::Continue);
}
for guardrail in output_guardrails {
let result = guardrail.check_dyn(&text).await;
if !result.is_pass() {
return Ok(map_guardrail_result(result, "output"));
}
}
Ok(HookAction::Continue)
}
_ => Ok(HookAction::Continue),
}
}
}
}