use std::collections::HashMap;
use crate::api::registry::{
EventMetadataInjector, ExecutionIntercept, Guardrail, Intercept, RuntimeRegistrationKind,
runtime_registration_is_enabled,
};
use crate::api::runtime::{
EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionFn, LlmRequestInterceptFn,
LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionFn, ToolConditionalFn,
ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn,
};
use crate::registry::SortedRegistry;
#[derive(Clone)]
pub(crate) struct ScopeLocalRegistries {
pub(crate) event_metadata_injectors: SortedRegistry<EventMetadataInjector>,
pub(crate) mark_sanitize_guardrails: SortedRegistry<Guardrail<EventSanitizeFn>>,
pub(crate) scope_sanitize_start_guardrails: SortedRegistry<Guardrail<EventSanitizeFn>>,
pub(crate) scope_sanitize_end_guardrails: SortedRegistry<Guardrail<EventSanitizeFn>>,
pub(crate) tool_sanitize_request_guardrails: SortedRegistry<Guardrail<ToolSanitizeFn>>,
pub(crate) tool_sanitize_response_guardrails: SortedRegistry<Guardrail<ToolSanitizeFn>>,
pub(crate) tool_conditional_execution_guardrails: SortedRegistry<Guardrail<ToolConditionalFn>>,
pub(crate) tool_request_intercepts: SortedRegistry<Intercept<ToolInterceptFn>>,
pub(crate) tool_execution_intercepts: SortedRegistry<ExecutionIntercept<ToolExecutionFn>>,
pub(crate) llm_sanitize_request_guardrails: SortedRegistry<Guardrail<LlmSanitizeRequestFn>>,
pub(crate) llm_sanitize_response_guardrails: SortedRegistry<Guardrail<LlmSanitizeResponseFn>>,
pub(crate) llm_conditional_execution_guardrails: SortedRegistry<Guardrail<LlmConditionalFn>>,
pub(crate) llm_request_intercepts: SortedRegistry<Intercept<LlmRequestInterceptFn>>,
pub(crate) llm_execution_intercepts: SortedRegistry<ExecutionIntercept<LlmExecutionFn>>,
pub(crate) llm_stream_execution_intercepts:
SortedRegistry<ExecutionIntercept<LlmStreamExecutionFn>>,
pub(crate) event_subscribers: HashMap<String, EventSubscriberFn>,
}
impl ScopeLocalRegistries {
pub(crate) fn new() -> Self {
Self {
event_metadata_injectors: SortedRegistry::new(),
mark_sanitize_guardrails: SortedRegistry::new(),
scope_sanitize_start_guardrails: SortedRegistry::new(),
scope_sanitize_end_guardrails: SortedRegistry::new(),
tool_sanitize_request_guardrails: SortedRegistry::new(),
tool_sanitize_response_guardrails: SortedRegistry::new(),
tool_conditional_execution_guardrails: SortedRegistry::new(),
tool_request_intercepts: SortedRegistry::new(),
tool_execution_intercepts: SortedRegistry::new(),
llm_sanitize_request_guardrails: SortedRegistry::new(),
llm_sanitize_response_guardrails: SortedRegistry::new(),
llm_conditional_execution_guardrails: SortedRegistry::new(),
llm_request_intercepts: SortedRegistry::new(),
llm_execution_intercepts: SortedRegistry::new(),
llm_stream_execution_intercepts: SortedRegistry::new(),
event_subscribers: HashMap::new(),
}
}
}
pub(crate) fn merge_event_metadata_injector_entries<'a>(
global: &'a SortedRegistry<EventMetadataInjector>,
scope_locals: &'a [&'a SortedRegistry<EventMetadataInjector>],
) -> Vec<&'a EventMetadataInjector> {
let mut all = Vec::new();
all.extend(global.sorted_values().into_iter().filter(|entry| {
runtime_registration_is_enabled(RuntimeRegistrationKind::EventMetadataInjector, &entry.name)
}));
for registry in scope_locals {
all.extend(registry.sorted_values());
}
all.sort_by(|left, right| {
left.priority
.cmp(&right.priority)
.then_with(|| left.name.cmp(&right.name))
});
all
}
impl Default for ScopeLocalRegistries {
fn default() -> Self {
Self::new()
}
}
pub(crate) fn merge_guardrail_entries<'a, F>(
global: &'a SortedRegistry<Guardrail<F>>,
scope_locals: &'a [&'a SortedRegistry<Guardrail<F>>],
kind: RuntimeRegistrationKind,
) -> Vec<&'a Guardrail<F>> {
let mut all = Vec::new();
all.extend(
global
.sorted_values()
.into_iter()
.filter(|entry| runtime_registration_is_enabled(kind, &entry.name)),
);
for registry in scope_locals {
all.extend(registry.sorted_values());
}
all.sort_by_key(|entry| entry.priority);
all
}
pub(crate) fn merge_intercept_entries<'a, F>(
global: &'a SortedRegistry<Intercept<F>>,
scope_locals: &'a [&'a SortedRegistry<Intercept<F>>],
kind: RuntimeRegistrationKind,
) -> Vec<&'a Intercept<F>> {
let mut all = Vec::new();
all.extend(
global
.sorted_values()
.into_iter()
.filter(|entry| runtime_registration_is_enabled(kind, &entry.name)),
);
for registry in scope_locals {
all.extend(registry.sorted_values());
}
all.sort_by_key(|entry| entry.priority);
all
}
pub(crate) fn merge_execution_intercept_callables<F: Clone>(
global: &SortedRegistry<ExecutionIntercept<F>>,
scope_locals: &[&SortedRegistry<ExecutionIntercept<F>>],
kind: RuntimeRegistrationKind,
) -> Vec<(F, i32)> {
let mut all = Vec::new();
for entry in global.sorted_values() {
if runtime_registration_is_enabled(kind, &entry.name) {
all.push((entry.payload.clone(), entry.priority));
}
}
for registry in scope_locals {
for entry in registry.sorted_values() {
all.push((entry.payload.clone(), entry.priority));
}
}
all.sort_by_key(|(_, priority)| *priority);
all
}