use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use futures_core::Stream;
use tower::Layer;
use tower::Service;
use crate::client::{BoxFuture, BoxStream};
use crate::error::{LiterLlmError, Result};
use crate::guardrail::registry::GuardrailRegistry;
use crate::guardrail::{GuardrailContext, GuardrailDecision, GuardrailStage};
use crate::types::ChatCompletionChunk;
use super::types::{LlmRequest, LlmRequestKind, LlmResponse};
fn serialize_for_guardrail<T: serde::Serialize>(value: &T) -> Result<serde_json::Value> {
serde_json::to_value(value).map_err(|e| LiterLlmError::InternalError {
message: format!("guardrail: failed to serialize response for output-stage inspection: {e}"),
})
}
fn request_to_guardrail_json(request: &LlmRequest) -> Result<serde_json::Value> {
serde_json::to_value(request).map_err(|e| LiterLlmError::InternalError {
message: format!("guardrail: failed to serialize request: {e}"),
})
}
pub const TENANT_ID_METADATA_KEY: &str = "tenant_id";
fn build_call_metadata(
layer_metadata: &Arc<HashMap<String, String>>,
request: &LlmRequest,
) -> Arc<HashMap<String, String>> {
let Some(tenant_id) = request.tenant_id() else {
return Arc::clone(layer_metadata);
};
if layer_metadata.contains_key(TENANT_ID_METADATA_KEY) {
tracing::warn!(
metadata_key = TENANT_ID_METADATA_KEY,
"guardrail: static per-layer metadata already defines this key; discarding the per-call value"
);
return Arc::clone(layer_metadata);
}
let mut merged = (**layer_metadata).clone();
merged.insert(TENANT_ID_METADATA_KEY.to_owned(), tenant_id.as_ref().to_owned());
Arc::new(merged)
}
fn apply_request_mutation(request: LlmRequest, new_payload: serde_json::Value) -> Result<LlmRequest> {
let mutated: LlmRequestKind = serde_json::from_value(new_payload).map_err(|e| LiterLlmError::InternalError {
message: format!("guardrail: Input stage Mutate payload is not a valid request: {e}"),
})?;
if std::mem::discriminant(&mutated) != std::mem::discriminant(&request.kind) {
return Err(LiterLlmError::InternalError {
message: "guardrail: Input stage Mutate payload changed the operation type".to_owned(),
});
}
Ok(LlmRequest {
kind: mutated,
tenant_id: request.tenant_id,
idempotency_key: request.idempotency_key,
})
}
fn apply_response_mutation(response: LlmResponse, new_payload: serde_json::Value) -> Result<LlmResponse> {
fn parse<T: serde::de::DeserializeOwned>(value: serde_json::Value) -> Result<T> {
serde_json::from_value(value).map_err(|e| LiterLlmError::InternalError {
message: format!("guardrail: Output stage Mutate payload is not a valid response: {e}"),
})
}
match response {
LlmResponse::Chat(_) => parse(new_payload).map(LlmResponse::Chat),
LlmResponse::Embed(_) => parse(new_payload).map(LlmResponse::Embed),
LlmResponse::ListModels(_) => parse(new_payload).map(LlmResponse::ListModels),
LlmResponse::ImageGenerate(_) => parse(new_payload).map(LlmResponse::ImageGenerate),
LlmResponse::Transcribe(_) => parse(new_payload).map(LlmResponse::Transcribe),
LlmResponse::Moderate(_) => parse(new_payload).map(LlmResponse::Moderate),
LlmResponse::Rerank(_) => parse(new_payload).map(LlmResponse::Rerank),
LlmResponse::Search(_) => parse(new_payload).map(LlmResponse::Search),
LlmResponse::Ocr(_) => parse(new_payload).map(LlmResponse::Ocr),
LlmResponse::Speech(_) | LlmResponse::ChatStream(_) => Err(LiterLlmError::InternalError {
message: "guardrail: Output stage Mutate is not supported for this response type".to_owned(),
}),
}
}
fn response_to_guardrail_json(response: &LlmResponse) -> Result<Option<serde_json::Value>> {
match response {
LlmResponse::Chat(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::Embed(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::ListModels(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::ImageGenerate(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::Transcribe(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::Moderate(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::Rerank(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::Search(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::Ocr(r) => serialize_for_guardrail(r).map(Some),
LlmResponse::Speech(audio_bytes) => Ok(Some(serde_json::json!({
"byte_len": audio_bytes.len(),
}))),
LlmResponse::ChatStream(_) => Ok(None),
}
}
fn chunk_text(chunk: &ChatCompletionChunk) -> String {
chunk
.choices
.iter()
.filter_map(|choice| choice.delta.content.as_deref())
.collect::<Vec<_>>()
.join("")
}
async fn apply_output_chunk_guardrail(
mut chunk: ChatCompletionChunk,
registry: &GuardrailRegistry,
request_json: &serde_json::Value,
metadata: &HashMap<String, String>,
) -> Result<ChatCompletionChunk> {
let text = chunk_text(&chunk);
if text.is_empty() {
return Ok(chunk);
}
let ctx = GuardrailContext {
request: request_json,
response: None,
chunk: Some(&text),
metadata,
};
match registry.run_stage(GuardrailStage::OutputChunk, &ctx).await {
GuardrailDecision::Block { reason, code } => Err(LiterLlmError::HookRejected {
message: format!("guardrail blocked output chunk [code={code}]: {reason}"),
}),
GuardrailDecision::Mutate { new_payload } => {
let replacement = new_payload.as_str().unwrap_or_default().to_owned();
for choice in &mut chunk.choices {
if choice.delta.content.is_some() {
choice.delta.content = Some(replacement.clone());
}
}
Ok(chunk)
}
GuardrailDecision::Allow => Ok(chunk),
}
}
struct GuardedChunkStream {
inner: BoxStream<'static, Result<ChatCompletionChunk>>,
registry: Arc<GuardrailRegistry>,
request_json: Arc<serde_json::Value>,
metadata: Arc<HashMap<String, String>>,
pending: Option<Pin<Box<dyn Future<Output = Result<ChatCompletionChunk>> + Send>>>,
blocked: bool,
}
impl Stream for GuardedChunkStream {
type Item = Result<ChatCompletionChunk>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if this.blocked {
return Poll::Ready(None);
}
loop {
if let Some(fut) = this.pending.as_mut() {
return match fut.as_mut().poll(cx) {
Poll::Ready(result) => {
this.pending = None;
if result.is_err() {
this.blocked = true;
}
Poll::Ready(Some(result))
}
Poll::Pending => Poll::Pending,
};
}
match this.inner.as_mut().poll_next(cx) {
Poll::Ready(Some(Ok(chunk))) => {
let registry = Arc::clone(&this.registry);
let request_json = Arc::clone(&this.request_json);
let metadata = Arc::clone(&this.metadata);
this.pending = Some(Box::pin(async move {
apply_output_chunk_guardrail(chunk, ®istry, &request_json, &metadata).await
}));
}
Poll::Ready(Some(Err(e))) => return Poll::Ready(Some(Err(e))),
Poll::Ready(None) => return Poll::Ready(None),
Poll::Pending => return Poll::Pending,
}
}
}
}
fn guard_output_chunk_stream(
stream: BoxStream<'static, Result<ChatCompletionChunk>>,
registry: Arc<GuardrailRegistry>,
request_json: Arc<serde_json::Value>,
metadata: Arc<HashMap<String, String>>,
) -> BoxStream<'static, Result<ChatCompletionChunk>> {
Box::pin(GuardedChunkStream {
inner: stream,
registry,
request_json,
metadata,
pending: None,
blocked: false,
})
}
#[cfg_attr(alef, alef(skip))]
#[derive(Clone)]
pub struct GuardrailLayer {
registry: Arc<GuardrailRegistry>,
metadata: Arc<HashMap<String, String>>,
}
impl GuardrailLayer {
#[must_use]
pub fn new(registry: Arc<GuardrailRegistry>, metadata: HashMap<String, String>) -> Self {
Self {
registry,
metadata: Arc::new(metadata),
}
}
#[must_use]
pub fn with_registry(registry: Arc<GuardrailRegistry>) -> Self {
Self::new(registry, HashMap::new())
}
}
impl<S> Layer<S> for GuardrailLayer {
type Service = GuardrailService<S>;
fn layer(&self, inner: S) -> Self::Service {
GuardrailService {
inner,
registry: Arc::clone(&self.registry),
metadata: Arc::clone(&self.metadata),
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct GuardrailService<S> {
inner: S,
registry: Arc<GuardrailRegistry>,
metadata: Arc<HashMap<String, String>>,
}
impl<S: Clone> Clone for GuardrailService<S> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
registry: Arc::clone(&self.registry),
metadata: Arc::clone(&self.metadata),
}
}
}
impl<S> Service<LlmRequest> for GuardrailService<S>
where
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Clone + Send + 'static,
S::Future: Send + 'static,
{
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<()>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: LlmRequest) -> Self::Future {
let registry = Arc::clone(&self.registry);
let metadata = if registry.is_empty() {
Arc::clone(&self.metadata)
} else {
build_call_metadata(&self.metadata, &req)
};
let standby = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, standby);
Box::pin(async move {
let request_json = Arc::new(request_to_guardrail_json(&req)?);
let input_ctx = GuardrailContext {
request: &request_json,
response: None,
chunk: None,
metadata: &metadata,
};
let input_decision = registry.run_stage(GuardrailStage::Input, &input_ctx).await;
let request_json = match input_decision {
GuardrailDecision::Block { reason, code } => {
return Err(LiterLlmError::HookRejected {
message: format!("guardrail blocked [code={code}]: {reason}"),
});
}
GuardrailDecision::Mutate { new_payload } => {
req = apply_request_mutation(req, new_payload)?;
Arc::new(request_to_guardrail_json(&req)?)
}
GuardrailDecision::Allow => request_json,
};
let response = inner.call(req).await?;
if let LlmResponse::ChatStream(stream) = response {
let guarded = guard_output_chunk_stream(
stream,
Arc::clone(®istry),
Arc::clone(&request_json),
Arc::clone(&metadata),
);
return Ok(LlmResponse::ChatStream(guarded));
}
let Some(response_json) = response_to_guardrail_json(&response)? else {
return Ok(response);
};
let output_ctx = GuardrailContext {
request: &request_json,
response: Some(&response_json),
chunk: None,
metadata: &metadata,
};
let output_decision = registry.run_stage(GuardrailStage::Output, &output_ctx).await;
match output_decision {
GuardrailDecision::Block { reason, code } => Err(LiterLlmError::HookRejected {
message: format!("guardrail blocked output [code={code}]: {reason}"),
}),
GuardrailDecision::Mutate { new_payload } => apply_response_mutation(response, new_payload),
GuardrailDecision::Allow => Ok(response),
}
})
}
}
#[cfg(test)]
mod tests {
use std::collections::{HashMap, HashSet};
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use tower::{Layer, Service};
use super::*;
use crate::guardrail::Guardrail;
use crate::guardrail::builtin::{AllowListGuardrail, DenyListGuardrail};
use crate::guardrail::registry::GuardrailRegistry;
use crate::tower::service::LlmService;
use crate::tower::tests_common::{MockClient, chat_req, make_chat_response};
use crate::tower::types::LlmRequest;
use crate::types::audio::{CreateSpeechRequest, CreateTranscriptionRequest, TranscriptionResponse};
use crate::types::common::{AssistantContent, Message, UserMessage};
use crate::types::image::{CreateImageRequest, ImagesResponse};
use crate::types::moderation::{ModerationRequest, ModerationResponse};
use crate::types::ocr::{OcrRequest, OcrResponse};
use crate::types::rerank::{RerankRequest, RerankResponse};
use crate::types::search::{SearchRequest, SearchResponse};
#[tokio::test]
async fn guardrail_layer_allows_when_registry_is_empty() {
let registry = Arc::new(GuardrailRegistry::new());
let inner = LlmService::new(MockClient::ok());
let mut svc = GuardrailLayer::with_registry(registry).layer(inner);
let result = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(result.is_ok(), "empty registry should allow all requests");
}
#[tokio::test]
async fn guardrail_layer_input_block_prevents_inner_call() {
let mut registry = GuardrailRegistry::new();
let list: HashSet<String> = ["banned-user"].iter().map(|s| s.to_string()).collect();
registry.register(Arc::new(DenyListGuardrail::new("ban", list, "user_id")));
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let mut meta = HashMap::new();
meta.insert("user_id".to_string(), "banned-user".to_string());
let mut svc = GuardrailLayer::new(Arc::new(registry), meta).layer(inner);
let err = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect_err("banned user should be blocked");
assert!(
matches!(err, LiterLlmError::HookRejected { .. }),
"guardrail block should surface as HookRejected"
);
assert_eq!(call_count.load(Ordering::SeqCst), 0, "inner service must not be called");
}
#[tokio::test]
async fn guardrail_layer_allows_non_blocked_user() {
let mut registry = GuardrailRegistry::new();
let list: HashSet<String> = ["banned-user"].iter().map(|s| s.to_string()).collect();
registry.register(Arc::new(DenyListGuardrail::new("ban", list, "user_id")));
let inner = LlmService::new(MockClient::ok());
let mut meta = HashMap::new();
meta.insert("user_id".to_string(), "good-user".to_string());
let mut svc = GuardrailLayer::new(Arc::new(registry), meta).layer(inner);
let result = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(result.is_ok(), "non-blocked user should pass through");
}
#[derive(Clone)]
struct CannedService {
build: Arc<dyn Fn() -> LlmResponse + Send + Sync>,
}
impl Service<LlmRequest> for CannedService {
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, _req: LlmRequest) -> Self::Future {
let resp = (self.build)();
Box::pin(async move { Ok(resp) })
}
}
struct AlwaysBlockOutput;
impl Guardrail for AlwaysBlockOutput {
fn name(&self) -> &'static str {
"always-block-output"
}
fn supported_stages(&self) -> &'static [GuardrailStage] {
&[GuardrailStage::Output]
}
fn check<'a>(
&'a self,
_stage: GuardrailStage,
_ctx: &'a GuardrailContext<'a>,
) -> Pin<Box<dyn std::future::Future<Output = GuardrailDecision> + Send + 'a>> {
Box::pin(async move {
GuardrailDecision::Block {
reason: "test: always blocks output".into(),
code: 9999,
}
})
}
}
async fn assert_output_stage_inspects<F>(request: LlmRequest, build_response: F)
where
F: Fn() -> LlmResponse + Send + Sync + 'static,
{
let mut registry = GuardrailRegistry::new();
registry.register(Arc::new(AlwaysBlockOutput));
let inner = CannedService {
build: Arc::new(build_response),
};
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let err = svc
.call(request)
.await
.expect_err("Output-stage guardrail should have blocked this response type");
assert!(
matches!(err, LiterLlmError::HookRejected { .. }),
"expected HookRejected, got {err:?}"
);
}
#[tokio::test]
async fn guardrail_output_stage_inspects_image_generate_response() {
assert_output_stage_inspects(LlmRequest::ImageGenerate(CreateImageRequest::default()), || {
LlmResponse::ImageGenerate(ImagesResponse::default())
})
.await;
}
#[tokio::test]
async fn guardrail_output_stage_inspects_speech_response() {
assert_output_stage_inspects(LlmRequest::Speech(CreateSpeechRequest::default()), || {
LlmResponse::Speech(bytes::Bytes::from_static(b"audio"))
})
.await;
}
#[tokio::test]
async fn guardrail_output_stage_inspects_transcribe_response() {
assert_output_stage_inspects(LlmRequest::Transcribe(CreateTranscriptionRequest::default()), || {
LlmResponse::Transcribe(TranscriptionResponse::default())
})
.await;
}
#[tokio::test]
async fn guardrail_output_stage_inspects_moderate_response() {
assert_output_stage_inspects(LlmRequest::Moderate(ModerationRequest::default()), || {
LlmResponse::Moderate(ModerationResponse {
id: String::new(),
model: String::new(),
results: vec![],
})
})
.await;
}
#[tokio::test]
async fn guardrail_output_stage_inspects_rerank_response() {
assert_output_stage_inspects(LlmRequest::Rerank(RerankRequest::default()), || {
LlmResponse::Rerank(RerankResponse {
id: None,
results: vec![],
meta: None,
})
})
.await;
}
#[tokio::test]
async fn guardrail_output_stage_inspects_search_response() {
assert_output_stage_inspects(LlmRequest::Search(SearchRequest::default()), || {
LlmResponse::Search(SearchResponse {
results: vec![],
model: "test-model".into(),
})
})
.await;
}
#[tokio::test]
async fn guardrail_output_stage_inspects_ocr_response() {
assert_output_stage_inspects(LlmRequest::Ocr(OcrRequest::default()), || {
LlmResponse::Ocr(OcrResponse {
pages: vec![],
model: "test-model".into(),
usage: None,
})
})
.await;
}
struct AlwaysFailsToSerialize;
impl serde::Serialize for AlwaysFailsToSerialize {
fn serialize<S: serde::Serializer>(&self, _serializer: S) -> std::result::Result<S::Ok, S::Error> {
Err(serde::ser::Error::custom("intentional failure for test"))
}
}
#[test]
fn serialize_for_guardrail_fails_closed_on_serialization_error() {
let result = serialize_for_guardrail(&AlwaysFailsToSerialize);
assert!(
result.is_err(),
"a response body that cannot be serialized must fail closed (Err), not silently pass through"
);
}
use crate::guardrail::builtin::{OnMatch, RegexGuardrail};
use crate::types::{ChatCompletionChunk, StreamChoice, StreamDelta};
use futures_util::StreamExt as _;
fn content_chunk(content: &str) -> Result<ChatCompletionChunk> {
Ok(ChatCompletionChunk {
id: "chunk".into(),
object: "chat.completion.chunk".into(),
created: 0,
model: "test-model".into(),
choices: vec![StreamChoice {
index: 0,
delta: StreamDelta {
content: Some(content.to_owned()),
..Default::default()
},
finish_reason: None,
}],
usage: None,
system_fingerprint: None,
service_tier: None,
})
}
struct VecChunkStream {
items: std::collections::VecDeque<Result<ChatCompletionChunk>>,
}
impl futures_core::Stream for VecChunkStream {
type Item = Result<ChatCompletionChunk>;
fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Ready(self.items.pop_front())
}
}
fn blocking_output_chunk_registry() -> GuardrailRegistry {
let mut registry = GuardrailRegistry::new();
static STAGES: &[GuardrailStage] = &[GuardrailStage::OutputChunk];
registry.register(Arc::new(RegexGuardrail::new(
"block-secret",
regex::Regex::new("SECRET").expect("valid regex"),
OnMatch::Block {
code: 1042,
reason_prefix: "secret leaked".into(),
},
STAGES,
)));
registry
}
#[tokio::test]
async fn guardrail_output_chunk_stage_blocks_streamed_phrase() {
let registry = blocking_output_chunk_registry();
let inner = CannedService {
build: Arc::new(|| {
let stream: crate::client::BoxStream<'static, Result<ChatCompletionChunk>> = Box::pin(VecChunkStream {
items: std::collections::VecDeque::from([
content_chunk("hello "),
content_chunk("this is SECRET data"),
]),
});
LlmResponse::ChatStream(stream)
}),
};
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let response = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("ChatStream response itself must not be rejected up front");
let LlmResponse::ChatStream(mut stream) = response else {
panic!("expected ChatStream response");
};
let first = stream.next().await.expect("first chunk must be yielded").expect(
"first chunk contains no blocked phrase and must pass through \
the OutputChunk guardrail unchanged",
);
assert_eq!(first.choices[0].delta.content.as_deref(), Some("hello "));
let second = stream.next().await.expect("second chunk must be yielded");
assert!(
matches!(second, Err(LiterLlmError::HookRejected { .. })),
"chunk containing the blocked phrase must surface as HookRejected, got {second:?}"
);
}
#[tokio::test]
async fn guardrail_output_chunk_stage_terminates_stream_after_block() {
let registry = blocking_output_chunk_registry();
let inner = CannedService {
build: Arc::new(|| {
let stream: crate::client::BoxStream<'static, Result<ChatCompletionChunk>> = Box::pin(VecChunkStream {
items: std::collections::VecDeque::from([
content_chunk("this is SECRET data"),
content_chunk("more content after the violation"),
]),
});
LlmResponse::ChatStream(stream)
}),
};
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let response = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("ChatStream response itself must not be rejected up front");
let LlmResponse::ChatStream(mut stream) = response else {
panic!("expected ChatStream response");
};
let first = stream
.next()
.await
.expect("blocked chunk must still be yielded once, as an Err");
assert!(matches!(first, Err(LiterLlmError::HookRejected { .. })));
let second = stream.next().await;
assert!(
second.is_none(),
"stream must terminate after a block, not yield the remaining queued chunk; got {second:?}"
);
}
#[tokio::test]
async fn guardrail_output_chunk_stage_mutate_redacts_and_continues() {
let mut registry = GuardrailRegistry::new();
static STAGES: &[GuardrailStage] = &[GuardrailStage::OutputChunk];
registry.register(Arc::new(RegexGuardrail::new(
"redact-secret",
regex::Regex::new("SECRET").expect("valid regex"),
OnMatch::Redact {
replacement: "[REDACTED]".into(),
},
STAGES,
)));
let inner = CannedService {
build: Arc::new(|| {
let stream: crate::client::BoxStream<'static, Result<ChatCompletionChunk>> = Box::pin(VecChunkStream {
items: std::collections::VecDeque::from([content_chunk("this is SECRET data")]),
});
LlmResponse::ChatStream(stream)
}),
};
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let response = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("call must succeed");
let LlmResponse::ChatStream(mut stream) = response else {
panic!("expected ChatStream response");
};
let first = stream
.next()
.await
.expect("chunk must be yielded")
.expect("mutate decision must not error");
assert_eq!(
first.choices[0].delta.content.as_deref(),
Some("this is [REDACTED] data"),
"matched text must be redacted in place"
);
assert!(stream.next().await.is_none(), "stream must end after the single chunk");
}
#[derive(Clone)]
struct RecordingService {
seen: Arc<std::sync::Mutex<Option<LlmRequest>>>,
}
impl Service<LlmRequest> for RecordingService {
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, req: LlmRequest) -> Self::Future {
*self.seen.lock().expect("lock") = Some(req);
Box::pin(async move { Ok(LlmResponse::Chat(make_chat_response("gpt-4"))) })
}
}
fn recorded_prompt(request: &LlmRequest) -> String {
let LlmRequestKind::Chat(chat) = &request.kind else {
panic!("expected a Chat request");
};
serde_json::to_string(&chat.messages).expect("messages must serialize")
}
#[tokio::test]
async fn guardrail_input_stage_mutate_rewrites_the_forwarded_request() {
let mut registry = GuardrailRegistry::new();
static STAGES: &[GuardrailStage] = &[GuardrailStage::Input];
registry.register(Arc::new(RegexGuardrail::new(
"redact-secret",
regex::Regex::new("SECRET").expect("valid regex"),
OnMatch::Redact {
replacement: "[REDACTED]".into(),
},
STAGES,
)));
let seen = Arc::new(std::sync::Mutex::new(None));
let inner = RecordingService {
seen: Arc::clone(&seen),
};
let mut chat = chat_req("gpt-4");
chat.messages = vec![Message::User(UserMessage {
content: "my password is SECRET".into(),
name: None,
})];
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
svc.call(LlmRequest::Chat(chat)).await.expect("call must succeed");
let forwarded = seen
.lock()
.expect("lock")
.clone()
.expect("inner service must be called");
let prompt = recorded_prompt(&forwarded);
assert!(
prompt.contains("[REDACTED]"),
"the mutated request must reach the inner service; got {prompt}"
);
assert!(
!prompt.contains("SECRET"),
"the original unredacted content must not reach the inner service; got {prompt}"
);
}
#[tokio::test]
async fn guardrail_input_stage_mutate_preserves_tenant_scope() {
let mut registry = GuardrailRegistry::new();
static STAGES: &[GuardrailStage] = &[GuardrailStage::Input];
registry.register(Arc::new(RegexGuardrail::new(
"redact-secret",
regex::Regex::new("SECRET").expect("valid regex"),
OnMatch::Redact {
replacement: "[REDACTED]".into(),
},
STAGES,
)));
let seen = Arc::new(std::sync::Mutex::new(None));
let inner = RecordingService {
seen: Arc::clone(&seen),
};
let mut chat = chat_req("gpt-4");
chat.messages = vec![Message::User(UserMessage {
content: "my password is SECRET".into(),
name: None,
})];
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
svc.call(
LlmRequest::Chat(chat)
.with_tenant_id("tenant-A")
.with_idempotency_key("idem-1"),
)
.await
.expect("call must succeed");
let forwarded = seen
.lock()
.expect("lock")
.clone()
.expect("inner service must be called");
assert_eq!(
forwarded.tenant_id().map(|t| t.as_ref().to_owned()),
Some("tenant-A".to_owned()),
"tenant must survive an Input-stage mutation"
);
assert_eq!(
forwarded.idempotency_key.as_deref(),
Some("idem-1"),
"idempotency key must survive an Input-stage mutation"
);
}
#[tokio::test]
async fn guardrail_output_stage_mutate_rewrites_the_returned_response() {
let mut registry = GuardrailRegistry::new();
static STAGES: &[GuardrailStage] = &[GuardrailStage::Output];
registry.register(Arc::new(RegexGuardrail::new(
"redact-secret",
regex::Regex::new("SECRET").expect("valid regex"),
OnMatch::Redact {
replacement: "[REDACTED]".into(),
},
STAGES,
)));
let inner = CannedService {
build: Arc::new(|| {
let mut resp = make_chat_response("gpt-4");
resp.choices[0].message.content = Some("the answer is SECRET".into());
LlmResponse::Chat(resp)
}),
};
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let response = svc
.call(LlmRequest::Chat(chat_req("gpt-4")))
.await
.expect("call must succeed");
let LlmResponse::Chat(chat) = response else {
panic!("expected a Chat response");
};
let Some(AssistantContent::Text(text)) = &chat.choices[0].message.content else {
panic!("expected text content on the returned response");
};
assert_eq!(
text, "the answer is [REDACTED]",
"the mutated response must be what the caller receives"
);
}
#[tokio::test]
async fn guardrail_input_stage_inapplicable_mutate_fails_closed() {
struct GarbageMutate;
impl Guardrail for GarbageMutate {
fn name(&self) -> &'static str {
"garbage-mutate"
}
fn supported_stages(&self) -> &'static [GuardrailStage] {
static STAGES: &[GuardrailStage] = &[GuardrailStage::Input];
STAGES
}
fn check<'a>(
&'a self,
_stage: GuardrailStage,
_ctx: &'a GuardrailContext<'a>,
) -> Pin<Box<dyn Future<Output = GuardrailDecision> + Send + 'a>> {
Box::pin(async {
GuardrailDecision::Mutate {
new_payload: serde_json::json!({ "NotAVariant": 1 }),
}
})
}
}
let mut registry = GuardrailRegistry::new();
registry.register(Arc::new(GarbageMutate));
let seen = Arc::new(std::sync::Mutex::new(None));
let inner = RecordingService {
seen: Arc::clone(&seen),
};
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let result = svc.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(result.is_err(), "an inapplicable Mutate must fail the call");
assert!(
seen.lock().expect("lock").is_none(),
"the original request must not be forwarded when the mutation cannot be applied"
);
}
struct RecordingMetadataGuardrail {
seen: Arc<std::sync::Mutex<Option<HashMap<String, String>>>>,
}
impl Guardrail for RecordingMetadataGuardrail {
fn name(&self) -> &'static str {
"recording-metadata"
}
fn supported_stages(&self) -> &'static [GuardrailStage] {
static STAGES: &[GuardrailStage] = &[GuardrailStage::Input];
STAGES
}
fn check<'a>(
&'a self,
_stage: GuardrailStage,
ctx: &'a GuardrailContext<'a>,
) -> Pin<Box<dyn Future<Output = GuardrailDecision> + Send + 'a>> {
let seen = Arc::clone(&self.seen);
let metadata = ctx.metadata.clone();
Box::pin(async move {
*seen.lock().expect("lock") = Some(metadata);
GuardrailDecision::Allow
})
}
}
#[tokio::test]
async fn deny_list_guardrail_blocks_request_whose_tenant_is_on_the_list() {
let mut registry = GuardrailRegistry::new();
let list: HashSet<String> = ["evil-tenant"].iter().map(|s| s.to_string()).collect();
registry.register(Arc::new(DenyListGuardrail::new("tenant-ban", list, "tenant_id")));
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let result = svc
.call(LlmRequest::Chat(chat_req("gpt-4")).with_tenant_id("evil-tenant"))
.await;
let err = result.expect_err("a tenant on the deny-list must be blocked");
assert!(
matches!(err, LiterLlmError::HookRejected { .. }),
"guardrail block should surface as HookRejected, got {err:?}"
);
assert_eq!(
call_count.load(Ordering::SeqCst),
0,
"inner service must not be called for a denied tenant"
);
}
#[tokio::test]
async fn deny_list_guardrail_allows_tenant_absent_from_list() {
let mut registry = GuardrailRegistry::new();
let list: HashSet<String> = ["evil-tenant"].iter().map(|s| s.to_string()).collect();
registry.register(Arc::new(DenyListGuardrail::new("tenant-ban", list, "tenant_id")));
let seen = Arc::new(std::sync::Mutex::new(None));
registry.register(Arc::new(RecordingMetadataGuardrail {
seen: Arc::clone(&seen),
}));
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let result = svc
.call(LlmRequest::Chat(chat_req("gpt-4")).with_tenant_id("good-tenant"))
.await;
assert!(
result.is_ok(),
"a tenant absent from the deny-list must be allowed through"
);
assert_eq!(
call_count.load(Ordering::SeqCst),
1,
"inner service must be called exactly once for an allowed tenant"
);
let recorded = seen
.lock()
.expect("lock")
.clone()
.expect("recording guardrail must have run");
assert_eq!(
recorded.get("tenant_id").map(String::as_str),
Some("good-tenant"),
"the per-call tenant_id must reach GuardrailContext::metadata; got {recorded:?}"
);
}
#[tokio::test]
async fn allow_list_guardrail_permits_listed_tenant() {
let mut registry = GuardrailRegistry::new();
let list: HashSet<String> = ["good-tenant"].iter().map(|s| s.to_string()).collect();
registry.register(Arc::new(AllowListGuardrail::new("tenant-allow", list, "tenant_id")));
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let result = svc
.call(LlmRequest::Chat(chat_req("gpt-4")).with_tenant_id("good-tenant"))
.await;
assert!(
result.is_ok(),
"a tenant on the allow-list must be permitted, got {result:?}"
);
assert_eq!(
call_count.load(Ordering::SeqCst),
1,
"inner service must be called once"
);
}
#[tokio::test]
async fn allow_list_guardrail_blocks_unlisted_tenant() {
let mut registry = GuardrailRegistry::new();
let list: HashSet<String> = ["good-tenant"].iter().map(|s| s.to_string()).collect();
registry.register(Arc::new(AllowListGuardrail::new("tenant-allow", list, "tenant_id")));
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let mut svc = GuardrailLayer::with_registry(Arc::new(registry)).layer(inner);
let err = svc
.call(LlmRequest::Chat(chat_req("gpt-4")).with_tenant_id("bad-tenant"))
.await
.expect_err("a tenant absent from the allow-list must be blocked");
let LiterLlmError::HookRejected { message } = err else {
panic!("expected HookRejected, got {err:?}");
};
assert!(
message.contains("code=1001") && message.contains("is not permitted"),
"block must be an evaluated value rejection, not a missing-field fail-closed; got {message}"
);
assert_eq!(call_count.load(Ordering::SeqCst), 0, "inner service must not be called");
}
#[tokio::test]
async fn static_layer_metadata_reaches_guardrail_alongside_per_call_tenant_id() {
let mut registry = GuardrailRegistry::new();
let seen = Arc::new(std::sync::Mutex::new(None));
registry.register(Arc::new(RecordingMetadataGuardrail {
seen: Arc::clone(&seen),
}));
let mut static_meta = HashMap::new();
static_meta.insert("route".to_string(), "prod-us-east".to_string());
let inner = LlmService::new(MockClient::ok());
let mut svc = GuardrailLayer::new(Arc::new(registry), static_meta).layer(inner);
svc.call(LlmRequest::Chat(chat_req("gpt-4")).with_tenant_id("tenant-A"))
.await
.expect("call must succeed");
let recorded = seen
.lock()
.expect("lock")
.clone()
.expect("recording guardrail must have run");
assert_eq!(
recorded.get("route").map(String::as_str),
Some("prod-us-east"),
"static per-layer metadata must survive the merge; got {recorded:?}"
);
assert_eq!(
recorded.get("tenant_id").map(String::as_str),
Some("tenant-A"),
"per-call tenant_id must be merged in alongside static metadata; got {recorded:?}"
);
}
#[tokio::test]
async fn static_metadata_wins_on_key_collision_with_per_call_tenant_id() {
let mut registry = GuardrailRegistry::new();
let seen = Arc::new(std::sync::Mutex::new(None));
registry.register(Arc::new(RecordingMetadataGuardrail {
seen: Arc::clone(&seen),
}));
let mut static_meta = HashMap::new();
static_meta.insert("tenant_id".to_string(), "static-tenant".to_string());
let inner = LlmService::new(MockClient::ok());
let mut svc = GuardrailLayer::new(Arc::new(registry), static_meta).layer(inner);
svc.call(LlmRequest::Chat(chat_req("gpt-4")).with_tenant_id("request-tenant"))
.await
.expect("call must succeed");
let recorded = seen
.lock()
.expect("lock")
.clone()
.expect("recording guardrail must have run");
assert_eq!(
recorded.get("tenant_id").map(String::as_str),
Some("static-tenant"),
"the static per-layer value must win on collision, not the per-call value; got {recorded:?}"
);
}
}