use std::sync::Arc;
use rpi_ai::event_stream::create_assistant_message_event_stream;
use rpi_ai::types::{AssistantMessage, DoneReason, ErrorReason};
use rpi_ai::{
AssistantMessageEvent, AssistantMessageEventStream, AssistantMessageEventStreamProducer,
Context, Model, Provider, SimpleStreamOptions,
};
use tokio::runtime::Handle;
use crate::loader::ExtensionSession;
use crate::registry::{RegisteredProvider, RegistrySnapshot};
use crate::PluginKeepalive;
pub struct PluggableProvider {
record: RegisteredProvider,
snapshot: Arc<RegistrySnapshot>,
_keepalive: Arc<PluginKeepalive>,
runtime: Handle,
}
impl PluggableProvider {
fn new(
record: RegisteredProvider,
snapshot: Arc<RegistrySnapshot>,
keepalive: Arc<PluginKeepalive>,
runtime: Handle,
) -> Arc<Self> {
Arc::new(Self {
record,
snapshot,
_keepalive: keepalive,
runtime,
})
}
pub fn from_session(session: &ExtensionSession, runtime: Handle) -> Vec<Arc<dyn Provider>> {
let Some(snapshot) = session.snapshot_arc() else {
return Vec::new();
};
let keepalive = session.keepalive();
snapshot
.providers()
.iter()
.cloned()
.map(|record| {
Self::new(
record,
Arc::clone(&snapshot),
Arc::clone(&keepalive),
runtime.clone(),
) as Arc<dyn Provider>
})
.collect()
}
}
#[async_trait::async_trait]
impl Provider for PluggableProvider {
fn id(&self) -> &str {
&self.record.provider_id
}
fn models(&self) -> &[Model] {
&[]
}
async fn stream_simple(
&self,
model: &Model,
ctx: &Context,
opts: &SimpleStreamOptions,
) -> AssistantMessageEventStream {
let (mut prod, stream) = create_assistant_message_event_stream();
let snapshot = Arc::clone(&self.snapshot);
let record = self.record.clone();
let runtime = self.runtime.clone();
let model = model.clone();
let ctx = Arc::new(ctx.clone());
let opts = opts.clone();
tokio::spawn(async move {
let message =
drive_plugin_provider(&snapshot, &record, runtime, &model, &ctx, &opts).await;
push_terminal(&mut prod, message);
});
stream
}
}
async fn drive_plugin_provider(
snapshot: &Arc<RegistrySnapshot>,
record: &RegisteredProvider,
runtime: Handle,
model: &Model,
ctx: &Context,
opts: &SimpleStreamOptions,
) -> AssistantMessage {
if !snapshot.is_active() {
return AssistantMessage::terminal(
model.api.clone(),
record.provider_id.clone(),
model.id.clone(),
rpi_ai::types::StopReason::Error,
"extensions provider registry is stale (session swapped/reloaded)",
0,
);
}
let request = serde_json::json!({
"model": serde_json::to_value(model).unwrap_or(serde_json::Value::Null),
"context": serde_json::to_value(ctx).unwrap_or(serde_json::Value::Null),
"options": options_json(opts),
});
let request_json = request.to_string();
let record = record.clone();
let provider_id = record.provider_id.clone();
let api = model.api.clone();
let model_id = model.id.clone();
let join = runtime.spawn_blocking(move || run_provider_request(&record, &request_json));
let outcome = match join.await {
Ok(inner) => inner,
Err(join_err) => {
return AssistantMessage::terminal(
api,
provider_id,
model_id,
rpi_ai::types::StopReason::Error,
format!("plugin provider task failed: {join_err}"),
0,
);
}
};
match outcome {
ProviderOutcome::Ok(message) => message,
ProviderOutcome::Err(msg) => AssistantMessage::terminal(
api,
provider_id,
model_id,
rpi_ai::types::StopReason::Error,
msg,
0,
),
}
}
enum ProviderOutcome {
Ok(AssistantMessage),
Err(String),
}
fn run_provider_request(record: &RegisteredProvider, request_json: &str) -> ProviderOutcome {
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
run_provider_request_inner(record, request_json)
}));
match outcome {
Ok(inner) => inner,
Err(_) => {
tracing::error!(
"plugin provider {} panicked — refusing to unwind across FFI",
record.provider_id
);
ProviderOutcome::Err(format!(
"plugin provider {} panicked during request",
record.provider_id
))
}
}
}
fn run_provider_request_inner(record: &RegisteredProvider, request_json: &str) -> ProviderOutcome {
let req_ref = StbStringRef::from_str(request_json);
let mut out = StbString::empty();
let rc = (record.request_fn)(req_ref, &mut out as *mut StbString, record.user_data);
if rc != 0 {
out.free_with(Some(record.plugin_free_string));
return ProviderOutcome::Err(format!(
"plugin provider {} returned error code {rc}",
record.provider_id
));
}
let response_text = out.to_string_lossy();
out.free_with(Some(record.plugin_free_string));
match serde_json::from_str::<AssistantMessage>(&response_text) {
Ok(message) => {
let mut message = message;
if message.provider.is_empty() {
message.provider = record.provider_id.clone();
}
if message.model.is_empty() {
message.model = String::new(); }
ProviderOutcome::Ok(message)
}
Err(err) => ProviderOutcome::Err(format!(
"plugin provider {} returned unparseable response: {err}",
record.provider_id
)),
}
}
fn push_terminal(prod: &mut AssistantMessageEventStreamProducer, message: AssistantMessage) {
use rpi_ai::types::StopReason;
match message.stop_reason {
StopReason::Error | StopReason::Aborted => {
prod.push(AssistantMessageEvent::Error {
reason: ErrorReason::Error,
error: message,
});
}
other => {
let reason = match other {
StopReason::Stop => DoneReason::Stop,
StopReason::Length => DoneReason::Length,
StopReason::ToolUse => DoneReason::ToolUse,
StopReason::Deferred => DoneReason::Deferred,
_ => DoneReason::Stop,
};
prod.push(AssistantMessageEvent::Done { reason, message });
}
}
}
fn options_json(opts: &SimpleStreamOptions) -> serde_json::Value {
serde_json::json!({
"apiKey": opts.api_key,
"timeoutMs": opts.timeout.map(|d| d.as_millis() as u64),
"maxRetries": opts.max_retries,
"maxRetryDelayMs": opts.max_retry_delay.map(|d| d.as_millis() as u64),
"headers": opts.headers,
"metadata": opts.metadata,
"cacheRetention": format!("{:?}", opts.cache_retention),
"sessionId": opts.session_id,
"reasoning": opts.reasoning.map(|r| format!("{r:?}")),
"maxTokens": opts.max_tokens,
"temperature": opts.temperature,
})
}
use rpi_plugin_sdk::{StbString, StbStringRef};
#[cfg(test)]
mod tests {
use super::*;
use crate::registry::ExtensionRegistry;
extern "C" fn ok_request_fn(
_req: StbStringRef,
out: *mut StbString,
_ud: *mut std::ffi::c_void,
) -> i32 {
let json = r#"{"role":"assistant","content":[{"type":"text","text":"hi"}],"api":"faux","provider":"pluggy","model":"m","usage":{"input":0,"output":0,"cacheRead":0,"cacheWrite":0,"totalTokens":0,"cost":{"input":0.0,"output":0.0,"cacheRead":0.0,"cacheWrite":0.0,"total":0.0}},"stopReason":"stop","timestamp":0}"#;
unsafe {
*out = StbString::from_string(json.to_string());
}
0
}
extern "C" fn stub_free(s: StbString) {
if s.is_empty() || s.ptr.is_null() {
return;
}
unsafe {
let slice = std::slice::from_raw_parts(s.ptr as *const u8, s.len);
let _ = Box::from_raw(slice as *const [u8] as *mut [u8]);
}
}
#[test]
fn request_fn_round_trip_ok_and_err() {
let mut registry = ExtensionRegistry::new();
registry.register_provider(RegisteredProvider {
provider_id: "pluggy".to_string(),
base_url: "https://example".to_string(),
api_style: "anthropic-messages".to_string(),
request_fn: ok_request_fn,
plugin_free_string: stub_free,
user_data: std::ptr::null_mut(),
});
let snap = Arc::new(registry.snapshot());
let record = snap.providers()[0].clone();
let outcome = run_provider_request(&record, r#"{"model":"m"}"#);
match outcome {
ProviderOutcome::Ok(m) => {
assert_eq!(m.provider, "pluggy");
assert_eq!(m.stop_reason, rpi_ai::types::StopReason::Stop);
assert_eq!(m.content.len(), 1);
}
_ => panic!("expected Ok"),
}
extern "C" fn err_request_fn(
_req: StbStringRef,
_out: *mut StbString,
_ud: *mut std::ffi::c_void,
) -> i32 {
7
}
let err_record = RegisteredProvider {
provider_id: "pluggy".to_string(),
base_url: "https://example".to_string(),
api_style: "anthropic-messages".to_string(),
request_fn: err_request_fn,
plugin_free_string: stub_free,
user_data: std::ptr::null_mut(),
};
let outcome = run_provider_request(&err_record, r#"{"model":"m"}"#);
assert!(matches!(outcome, ProviderOutcome::Err(_)));
}
}