use std::collections::HashMap;
use std::sync::{Arc, Mutex, OnceLock};
use openai_frontend::{OpenAiError, OpenAiResult};
pub trait GenerationGate: Send + Sync {
fn after_prefill(&self, input_tokens: usize, max_output_tokens: u32) -> OpenAiResult<()>;
fn before_token(&self) -> OpenAiResult<()> {
Ok(())
}
fn committed_token(&self) -> OpenAiResult<()>;
fn committed_tokens(&self) -> u64 {
0
}
}
type Gates = HashMap<[u8; 16], Arc<dyn GenerationGate>>;
fn registry() -> &'static Mutex<Gates> {
static REGISTRY: OnceLock<Mutex<Gates>> = OnceLock::new();
REGISTRY.get_or_init(Mutex::default)
}
pub struct Registration([u8; 16]);
impl Drop for Registration {
fn drop(&mut self) {
if let Ok(mut gates) = registry().lock() {
gates.remove(&self.0);
}
}
}
pub fn register(id: [u8; 16], gate: Arc<dyn GenerationGate>) -> OpenAiResult<Registration> {
let mut gates = registry()
.lock()
.map_err(|_| OpenAiError::backend("generation gate registry unavailable"))?;
if gates.contains_key(&id) {
return Err(OpenAiError::backend("duplicate generation gate"));
}
gates.insert(id, gate);
Ok(Registration(id))
}
pub(crate) fn find(id: Option<[u8; 16]>) -> OpenAiResult<Option<Arc<dyn GenerationGate>>> {
let gates = registry()
.lock()
.map_err(|_| OpenAiError::backend("generation gate registry unavailable"))?;
Ok(id.and_then(|id| gates.get(&id).cloned()))
}
#[cfg(feature = "test-support")]
pub fn registered_for_test(id: [u8; 16]) -> OpenAiResult<Option<Arc<dyn GenerationGate>>> {
find(Some(id))
}
pub(in crate::frontend) fn usage_event(
gate: Option<&Arc<dyn GenerationGate>>,
) -> Option<crate::frontend::generation::GenerationStreamEvent> {
gate.map(|gate| {
crate::frontend::generation::GenerationStreamEvent::Usage(
openai_frontend::Usage::new(0, gate.committed_tokens().min(u64::from(u32::MAX)) as u32),
None,
)
})
}