use std::ffi::c_void;
use std::panic::{catch_unwind, AssertUnwindSafe};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use async_trait::async_trait;
use rpi_agent::agent_tool::AgentTool;
use rpi_agent::error::AgentError;
use rpi_agent::types::{AgentToolResult, TextContentOrImage, ToolExecutionMode, ToolResultPartial};
use rpi_ai::types::Tool;
use rpi_plugin_sdk::{
FreeStringFn, StbString, StbStringRef, StepHandle, StepResultTag, ToolCancelFn, ToolDestroyFn,
ToolExecuteFn, ToolPollFn,
};
use tokio::sync::{mpsc, oneshot};
use tokio_util::sync::CancellationToken;
use crate::host_free_string;
use crate::loader::PluginKeepalive;
#[derive(Clone, Copy)]
pub struct PluginToolHandle {
pub(crate) execute_fn: ToolExecuteFn,
pub(crate) poll_fn: ToolPollFn,
pub(crate) cancel_fn: ToolCancelFn,
pub(crate) destroy_fn: ToolDestroyFn,
pub(crate) plugin_free_string: FreeStringFn,
}
fn result_to_json(result: &AgentToolResult) -> String {
let mut txt = String::new();
txt.push('{');
txt.push_str("\"content\":[");
for (i, c) in result.content.iter().enumerate() {
if i > 0 {
txt.push(',');
}
match c {
TextContentOrImage::Text(t) => {
txt.push_str(
&serde_json::to_string(&serde_json::json!({ "type": "text", "text": t.text }))
.unwrap_or_else(|_| "\"\"".into()),
);
}
TextContentOrImage::Image(img) => {
txt.push_str(
&serde_json::to_string(&serde_json::json!({
"type": "image",
"data": img.data,
"mimeType": img.mime_type,
}))
.unwrap_or_else(|_| "\"\"".into()),
);
}
}
}
txt.push(']');
txt.push_str(",\"details\":");
txt.push_str(&serde_json::to_string(&result.details).unwrap_or_else(|_| "null".into()));
txt.push_str(",\"terminate\":");
txt.push_str(if result.terminate { "true" } else { "false" });
txt.push_str(",\"addedToolNames\":");
txt.push_str(&serde_json::to_string(&result.added_tool_names).unwrap_or_else(|_| "[]".into()));
txt.push('}');
txt
}
fn stb_to_result(s: &StbString) -> AgentToolResult {
let text = s.to_string_lossy();
let val: serde_json::Value = serde_json::from_str(&text).unwrap_or(serde_json::Value::Null);
let mut result = AgentToolResult::default();
if let Some(obj) = val.as_object() {
if let Some(content) = obj.get("content").and_then(|v| v.as_array()) {
for block in content {
let kind = block.get("type").and_then(|v| v.as_str()).unwrap_or("text");
match kind {
"image" => {
let data = block
.get("data")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let mime = block
.get("mimeType")
.or_else(|| block.get("mime_type"))
.and_then(|v| v.as_str())
.unwrap_or("image/png")
.to_string();
result.content.push(TextContentOrImage::Image(
rpi_ai::types::ImageContent {
kind: rpi_ai::types::ImageContentType,
data,
mime_type: mime,
},
));
}
_ => {
let t = block
.get("text")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
result.content.push(TextContentOrImage::text(t));
}
}
}
}
if let Some(details) = obj.get("details") {
result.details = details.clone();
}
if let Some(terms) = obj.get("terminate").and_then(|v| v.as_bool()) {
result.terminate = terms;
}
if let Some(arr) = obj
.get("addedToolNames")
.or_else(|| obj.get("added_tool_names"))
.and_then(|v| v.as_array())
{
result.added_tool_names = arr
.iter()
.filter_map(|v| v.as_str().map(String::from))
.collect();
}
if let Some(usage) = obj.get("usage") {
if let Ok(u) = serde_json::from_value::<rpi_ai::types::Usage>(usage.clone()) {
result.usage = Some(u);
}
}
}
result
}
pub struct PluginToolAdapter {
schema: Tool,
label: String,
handle: PluginToolHandle,
#[allow(dead_code)]
keepalive: Arc<PluginKeepalive>,
}
impl PluginToolAdapter {
pub fn new(schema: Tool, handle: PluginToolHandle, keepalive: Arc<PluginKeepalive>) -> Self {
let label = schema.name.clone();
Self {
schema,
label,
handle,
keepalive,
}
}
}
extern "C" fn partial_cb_trampoline(partial: StbString, user_data: *mut c_void) {
let outcome = catch_unwind(AssertUnwindSafe(|| {
if user_data.is_null() {
host_free_string(partial);
return;
}
let sender = unsafe { &*(user_data as *const mpsc::UnboundedSender<AgentToolResult>) };
let result = stb_to_result(&partial);
host_free_string(partial);
let _ = sender.send(result);
}));
if outcome.is_err() {
tracing::error!("plugin partial callback panicked — aborting (cannot unwind across FFI)");
std::process::abort();
}
}
#[async_trait]
impl AgentTool for PluginToolAdapter {
fn schema(&self) -> &Tool {
&self.schema
}
fn label(&self) -> &str {
&self.label
}
fn execution_mode(&self) -> ToolExecutionMode {
ToolExecutionMode::Parallel
}
async fn execute(
&self,
tool_call_id: &str,
params: serde_json::Value,
signal: CancellationToken,
on_update: Arc<dyn Fn(ToolResultPartial) + Send + Sync>,
) -> Result<AgentToolResult, AgentError> {
let runtime = tokio::runtime::Handle::try_current().map_err(|e| {
AgentError::State(format!(
"plugin tool '{}' executed off-runtime: {e}",
self.schema.name
))
})?;
let (partial_tx, mut partial_rx) = mpsc::unbounded_channel::<AgentToolResult>();
let (done_tx, done_rx) = oneshot::channel::<Result<AgentToolResult, AgentError>>();
let cancel_flag = Arc::new(AtomicBool::new(false));
let params_json = serde_json::to_string(¶ms).unwrap_or_else(|_| "null".to_string());
let params_stb = StbString::from_string(params_json);
let id_ref = StbStringRef::from_str(tool_call_id);
let plugin_free = self.handle.plugin_free_string;
let execute_fn = self.handle.execute_fn;
let poll_fn = self.handle.poll_fn;
let cancel_fn = self.handle.cancel_fn;
let destroy_fn = self.handle.destroy_fn;
let schema_name = self.schema.name.clone();
let schema_name_for_error = schema_name.clone();
let cancel_flag_drive = Arc::clone(&cancel_flag);
let sender_for_cb = partial_tx.clone();
runtime.spawn_blocking(move || {
let step_handle: StepHandle = {
let outcome = catch_unwind(AssertUnwindSafe(|| {
(execute_fn)(id_ref, params_stb, Some(host_free_string))
}));
match outcome {
Ok(h) if !h.is_null() => h,
Ok(_) => {
let _ = done_tx.send(Err(AgentError::Tool(format!(
"plugin execute returned null handle for '{schema_name}'"
))));
return;
}
Err(_) => {
tracing::error!("plugin execute panicked — aborting");
std::process::abort();
}
}
};
let sender_ptr = &sender_for_cb as *const _ as *mut c_void;
let terminal: Result<AgentToolResult, AgentError> = loop {
if cancel_flag_drive.load(Ordering::SeqCst) {
let _ = catch_unwind(AssertUnwindSafe(|| (cancel_fn)(step_handle)));
break Err(AgentError::Tool("plugin tool cancelled".into()));
}
let step_result = match catch_unwind(AssertUnwindSafe(|| {
(poll_fn)(step_handle, Some(partial_cb_trampoline), sender_ptr)
})) {
Ok(r) => r,
Err(_) => {
tracing::error!("plugin poll panicked — aborting");
std::process::abort();
}
};
match step_result.tag {
StepResultTag::Pending => {
let progress = unsafe { step_result.pending_payload().progress };
if !progress.is_empty() {
let pr = stb_to_result(&progress);
(plugin_free)(progress);
let _ = partial_tx.send(pr);
}
continue;
}
StepResultTag::Done => {
let done = unsafe { step_result.done_payload().result };
let result = stb_to_result(&done);
(plugin_free)(done);
break Ok(result);
}
StepResultTag::Err => {
let msg = unsafe { step_result.err_payload().message };
let message = msg.to_string_lossy();
(plugin_free)(msg);
break Err(AgentError::Tool(message));
}
}
};
let _ = catch_unwind(AssertUnwindSafe(|| (destroy_fn)(step_handle)));
let _ = done_tx.send(terminal);
});
let schema_name_err = schema_name_for_error.clone();
tokio::pin!(done_rx);
let mut cancelled = std::pin::pin!(signal.cancelled());
let result: Result<AgentToolResult, AgentError> = loop {
tokio::select! {
done = &mut done_rx => {
while let Ok(p) = partial_rx.try_recv() { on_update(p); }
break match done {
Ok(Ok(r)) => Ok(r),
Ok(Err(e)) => Err(e),
Err(_) => Err(AgentError::State(format!(
"plugin tool '{schema_name_err}' driver dropped done_tx"
))),
};
}
_ = &mut cancelled => {
let already = cancel_flag.swap(true, Ordering::SeqCst);
if !already {
while let Ok(p) = partial_rx.try_recv() { on_update(p); }
}
tokio::task::yield_now().await;
}
}
};
while let Ok(p) = partial_rx.try_recv() {
on_update(p);
}
result
}
}
#[cfg(test)]
mod tests {
use super::*;
use rpi_plugin_sdk::{StepResult, ToolPartialCb};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Mutex;
static TEST_LOCK: Mutex<()> = Mutex::new(());
static DESTROY_COUNT: AtomicUsize = AtomicUsize::new(0);
static CANCEL_COUNT: AtomicUsize = AtomicUsize::new(0);
struct DriveState {
cancelled: Arc<AtomicBool>,
polls: usize,
done_at: usize,
}
extern "C" fn stub_execute(
_id: StbStringRef,
_params: StbString,
_free: Option<FreeStringFn>,
) -> StepHandle {
let state = Box::new(DriveState {
cancelled: Arc::new(AtomicBool::new(false)),
polls: 0,
done_at: 3,
});
Box::into_raw(state) as StepHandle
}
extern "C" fn stub_poll(
h: StepHandle,
_cb: Option<ToolPartialCb>,
_ud: *mut c_void,
) -> StepResult {
let state = unsafe { &mut *(h as *mut DriveState) };
state.polls += 1;
if state.cancelled.load(Ordering::SeqCst) {
return StepResult::err(StbString::from_string("cancelled".into()));
}
if state.polls >= state.done_at {
let result_json = StbString::from_string(
r#"{"content":[{"type":"text","text":"echo: hello"}]}"#.to_string(),
);
StepResult::done(result_json)
} else {
StepResult::pending(StbString::from_string(
r#"{"content":[{"type":"text","text":"..."}]}"#.to_string(),
))
}
}
extern "C" fn stub_cancel(h: StepHandle) {
CANCEL_COUNT.fetch_add(1, Ordering::SeqCst);
let state = unsafe { &*(h as *const DriveState) };
state.cancelled.store(true, Ordering::SeqCst);
}
extern "C" fn stub_destroy(h: StepHandle) {
DESTROY_COUNT.fetch_add(1, Ordering::SeqCst);
if h.is_null() {
return;
}
unsafe {
let _ = Box::from_raw(h as *mut DriveState);
}
}
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]);
}
}
fn stub_handle() -> PluginToolHandle {
PluginToolHandle {
execute_fn: stub_execute,
poll_fn: stub_poll,
cancel_fn: stub_cancel,
destroy_fn: stub_destroy,
plugin_free_string: stub_free,
}
}
fn echo_adapter() -> PluginToolAdapter {
let tool = Tool {
name: "echo".to_string(),
description: "echoes".to_string(),
parameters: rpi_ai::types::Schema::new(serde_json::json!({})),
constrained_sampling: None,
};
PluginToolAdapter::new(tool, stub_handle(), PluginKeepalive::empty())
}
fn reset_counters() {
DESTROY_COUNT.store(0, Ordering::SeqCst);
CANCEL_COUNT.store(0, Ordering::SeqCst);
}
#[tokio::test]
async fn adapter_drives_to_done_and_destroys_once() {
let _guard = TEST_LOCK.lock().unwrap();
reset_counters();
let adapter = echo_adapter();
let on_update: Arc<dyn Fn(ToolResultPartial) + Send + Sync> = Arc::new(|_| {});
let signal = CancellationToken::new();
let result = adapter
.execute("call_1", serde_json::json!({}), signal, on_update)
.await
.expect("drive should succeed");
assert_eq!(result.content.len(), 1);
assert_eq!(
DESTROY_COUNT.load(Ordering::SeqCst),
1,
"destroy exactly once"
);
assert_eq!(
CANCEL_COUNT.load(Ordering::SeqCst),
0,
"no cancel in happy path"
);
}
#[tokio::test]
async fn adapter_forwards_partials_to_on_update() {
let _guard = TEST_LOCK.lock().unwrap();
reset_counters();
let adapter = echo_adapter();
let seen = Arc::new(Mutex::new(Vec::<String>::new()));
let seen_clone = Arc::clone(&seen);
let on_update: Arc<dyn Fn(ToolResultPartial) + Send + Sync> = Arc::new(move |p| {
if let Some(t) = p.content.first().and_then(|c| match c {
TextContentOrImage::Text(t) => Some(t.text.clone()),
_ => None,
}) {
seen_clone.lock().unwrap().push(t);
}
});
let signal = CancellationToken::new();
let _ = adapter
.execute("call_2", serde_json::json!({}), signal, on_update)
.await
.expect("ok");
let partials = seen.lock().unwrap().clone();
assert!(partials.iter().any(|t| t == "..."), "got {:?}", partials);
assert_eq!(DESTROY_COUNT.load(Ordering::SeqCst), 1);
}
#[tokio::test(flavor = "current_thread")]
async fn adapter_cancel_observed_no_uaf_no_leak() {
let _guard = TEST_LOCK.lock().unwrap();
reset_counters();
struct SlowState {
cancelled: Arc<AtomicBool>,
polls: usize,
}
extern "C" fn slow_execute(
_: StbStringRef,
_: StbString,
_: Option<FreeStringFn>,
) -> StepHandle {
Box::into_raw(Box::new(SlowState {
cancelled: Arc::new(AtomicBool::new(false)),
polls: 0,
})) as StepHandle
}
extern "C" fn slow_poll(
h: StepHandle,
_: Option<ToolPartialCb>,
_: *mut c_void,
) -> StepResult {
let s = unsafe { &mut *(h as *mut SlowState) };
s.polls += 1;
if s.cancelled.load(Ordering::SeqCst) {
return StepResult::err(StbString::from_string("cancelled".into()));
}
std::thread::sleep(std::time::Duration::from_millis(5));
StepResult::pending(StbString::empty())
}
extern "C" fn slow_cancel(h: StepHandle) {
CANCEL_COUNT.fetch_add(1, Ordering::SeqCst);
unsafe {
(*(h as *mut SlowState))
.cancelled
.store(true, Ordering::SeqCst);
}
}
extern "C" fn slow_destroy(h: StepHandle) {
DESTROY_COUNT.fetch_add(1, Ordering::SeqCst);
if !h.is_null() {
unsafe {
let _ = Box::from_raw(h as *mut SlowState);
}
}
}
let tool = Tool {
name: "slow".to_string(),
description: "slow".to_string(),
parameters: rpi_ai::types::Schema::new(serde_json::json!({})),
constrained_sampling: None,
};
let handle = PluginToolHandle {
execute_fn: slow_execute,
poll_fn: slow_poll,
cancel_fn: slow_cancel,
destroy_fn: slow_destroy,
plugin_free_string: stub_free,
};
let adapter = PluginToolAdapter::new(tool, handle, PluginKeepalive::empty());
let on_update: Arc<dyn Fn(ToolResultPartial) + Send + Sync> = Arc::new(|_| {});
let signal = CancellationToken::new();
let signal_clone = signal.clone();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
signal_clone.cancel();
});
let _ = adapter
.execute("call_3", serde_json::json!({}), signal, on_update)
.await;
assert_eq!(
DESTROY_COUNT.load(Ordering::SeqCst),
1,
"destroy once on cancel"
);
}
}