use std::sync::Arc;
use rpi_ai::types::AssistantMessage;
use rpi_ai::Context;
use rpi_ai::{Model, ProviderHooks, SimpleStreamOptions, SimpleStreamOptionsPatch};
use rpi_plugin_sdk::EventTag;
use crate::loader::ExtensionSession;
use crate::registry::RegistrySnapshot;
use crate::translate::dispatch_data_event;
use crate::PluginKeepalive;
pub struct ExtensionProviderHooks {
snapshot: Arc<RegistrySnapshot>,
_keepalive: Arc<PluginKeepalive>,
}
impl ExtensionProviderHooks {
pub fn from_session(session: &ExtensionSession) -> Option<Self> {
let snapshot = session.snapshot_arc()?;
let subscribed = [
EventTag::BeforeProviderRequest,
EventTag::BeforeProviderHeaders,
EventTag::AfterProviderResponse,
]
.iter()
.any(|t| !snapshot.handlers_for(*t).is_empty());
if !subscribed {
return None;
}
Some(Self {
snapshot,
_keepalive: session.keepalive(),
})
}
}
impl ProviderHooks for ExtensionProviderHooks {
fn before_request(
&self,
model: &Model,
_ctx: &Context,
opts: &SimpleStreamOptions,
) -> Option<SimpleStreamOptionsPatch> {
let request = serde_json::json!({
"model": model.id,
"provider": model.provider,
"baseUrl": model.base_url,
"reasoning": model.reasoning,
"apiKey": opts.api_key,
"timeoutMs": opts.timeout.map(|d| d.as_millis() as u64),
"headers": opts.headers,
"metadata": opts.metadata,
"maxTokens": opts.max_tokens,
"temperature": opts.temperature,
});
dispatch_data_event(
&self.snapshot,
EventTag::BeforeProviderRequest,
&request.to_string(),
);
let headers = serde_json::json!({ "headers": opts.headers });
dispatch_data_event(
&self.snapshot,
EventTag::BeforeProviderHeaders,
&headers.to_string(),
);
None
}
fn after_response(&self, _model: &Model, message: &AssistantMessage) {
let json = serde_json::to_value(message).unwrap_or(serde_json::json!({}));
dispatch_data_event(
&self.snapshot,
EventTag::AfterProviderResponse,
&json.to_string(),
);
}
}
#[cfg(test)]
mod tests {
use super::*;
use rpi_plugin_sdk::{EventTag, StablePluginEvent};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Mutex;
static PROVIDER_HITS: AtomicUsize = AtomicUsize::new(0);
static PROVIDER_LOCK: Mutex<()> = Mutex::new(());
extern "C" fn counting_provider_handler(
ev: StablePluginEvent,
_ud: *mut std::ffi::c_void,
) -> i32 {
let s = unsafe { ev.payload.data.data.to_string_lossy() };
if let Ok(v) = serde_json::from_str::<serde_json::Value>(&s) {
if v.get("model").is_some() {
assert!(v.get("model").is_some(), "request event must carry model");
}
PROVIDER_HITS.fetch_add(1, Ordering::SeqCst);
}
0
}
#[test]
fn provider_hooks_dispatch_to_subscribers() {
let _guard = PROVIDER_LOCK.lock().unwrap();
PROVIDER_HITS.store(0, Ordering::SeqCst);
let mut registry = crate::registry::ExtensionRegistry::new();
let handler: rpi_plugin_sdk::EventHandlerFn = counting_provider_handler;
registry.register_event_handler(
EventTag::BeforeProviderRequest,
handler,
std::ptr::null_mut(),
);
registry.register_event_handler(
EventTag::BeforeProviderHeaders,
handler,
std::ptr::null_mut(),
);
let snapshot = Arc::new(registry.snapshot());
assert!(dispatch_data_event(
&snapshot,
EventTag::BeforeProviderRequest,
r#"{"model":"m"}"#,
));
assert_eq!(PROVIDER_HITS.load(Ordering::SeqCst), 1);
let hooks = ExtensionProviderHooks {
snapshot: Arc::clone(&snapshot),
_keepalive: Arc::new(crate::PluginKeepalive::new(Vec::new(), None)),
};
let model = rpi_ai::Model::new(
"m",
"m",
rpi_ai::Api::AnthropicMessages,
"anthropic",
"https://api.anthropic.com",
);
let ctx = rpi_ai::Context::default();
let opts = rpi_ai::SimpleStreamOptions::default();
let patch = hooks.before_request(&model, &ctx, &opts);
assert!(patch.is_none(), "observer-only v1: no patch returned");
assert_eq!(PROVIDER_HITS.load(Ordering::SeqCst), 3);
}
#[test]
fn provider_hooks_no_subscribers_is_noop() {
let registry = crate::registry::ExtensionRegistry::new();
let snapshot = Arc::new(registry.snapshot());
assert!(!dispatch_data_event(
&snapshot,
EventTag::AfterProviderResponse,
"{}",
));
}
}