use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use rpi_agent::error::AgentError;
use rpi_ai::types::Tool;
use rpi_plugin_sdk::{
CommandHandlerFn, EventHandlerFn, EventTag, FreeStringFn, ProviderRequestFn, RenderFn,
ResourcesDiscoverFn, EVENT_TAG_COUNT,
};
use crate::tool::PluginToolHandle;
#[derive(Clone)]
pub struct ExtensionTool {
pub tool: Tool,
pub(crate) handle: PluginToolHandle,
}
impl ExtensionTool {
pub fn new(tool: Tool, handle: PluginToolHandle) -> Self {
Self { tool, handle }
}
pub fn handle(&self) -> PluginToolHandle {
self.handle
}
}
#[derive(Clone)]
pub struct RegisteredCommand {
pub name: String,
pub description: String,
pub handler: CommandHandlerFn,
pub user_data: *mut std::ffi::c_void,
}
unsafe impl Send for RegisteredCommand {}
unsafe impl Sync for RegisteredCommand {}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RegisteredFlag {
pub name: String,
pub description: String,
}
#[derive(Clone, Copy)]
pub struct RegisteredHandler {
pub handler: EventHandlerFn,
pub user_data: *mut std::ffi::c_void,
}
unsafe impl Send for RegisteredHandler {}
unsafe impl Sync for RegisteredHandler {}
#[derive(Clone, Copy)]
pub struct ResourcesDiscoverHandler {
pub handler: ResourcesDiscoverFn,
pub plugin_free_string: FreeStringFn,
pub user_data: *mut std::ffi::c_void,
}
unsafe impl Send for ResourcesDiscoverHandler {}
unsafe impl Sync for ResourcesDiscoverHandler {}
#[derive(Clone)]
pub struct RegisteredProvider {
pub provider_id: String,
pub base_url: String,
pub api_style: String,
pub request_fn: ProviderRequestFn,
pub plugin_free_string: FreeStringFn,
pub user_data: *mut std::ffi::c_void,
}
unsafe impl Send for RegisteredProvider {}
unsafe impl Sync for RegisteredProvider {}
#[derive(Clone)]
pub struct RegisteredRenderer {
pub name: String,
pub kind: RegisteredRendererKind,
pub render_fn: RenderFn,
pub plugin_free_string: FreeStringFn,
pub user_data: *mut std::ffi::c_void,
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub enum RegisteredRendererKind {
Message,
Markdown,
Entry,
}
unsafe impl Send for RegisteredRenderer {}
unsafe impl Sync for RegisteredRenderer {}
#[derive(Clone)]
pub enum RegistryEntry {
Tool(ExtensionTool),
Command(RegisteredCommand),
Flag(RegisteredFlag),
}
pub struct ExtensionRegistry {
tools: Vec<ExtensionTool>,
commands: Vec<RegisteredCommand>,
flags: Vec<RegisteredFlag>,
handlers: [Vec<RegisteredHandler>; EVENT_TAG_COUNT],
resources_discover: Vec<ResourcesDiscoverHandler>,
providers: Vec<RegisteredProvider>,
renderers: Vec<RegisteredRenderer>,
active: Arc<AtomicBool>,
}
impl Default for ExtensionRegistry {
fn default() -> Self {
Self::new()
}
}
impl ExtensionRegistry {
pub fn new() -> Self {
let handlers: [Vec<RegisteredHandler>; EVENT_TAG_COUNT] =
std::array::from_fn(|_| Vec::new());
Self {
tools: Vec::new(),
commands: Vec::new(),
flags: Vec::new(),
handlers,
resources_discover: Vec::new(),
providers: Vec::new(),
renderers: Vec::new(),
active: Arc::new(AtomicBool::new(true)),
}
}
pub fn register_tool(&mut self, tool: Tool, handle: PluginToolHandle) -> bool {
if self.tools.iter().any(|t| t.tool.name == tool.name) {
return true;
}
self.tools.push(ExtensionTool::new(tool, handle));
false
}
pub fn register_command(
&mut self,
name: String,
description: String,
handler: CommandHandlerFn,
user_data: *mut std::ffi::c_void,
) -> bool {
if self.commands.iter().any(|c| c.name == name) {
return true;
}
self.commands.push(RegisteredCommand {
name,
description,
handler,
user_data,
});
false
}
pub fn register_flag(&mut self, name: String, description: String) -> bool {
if self.flags.iter().any(|flag| flag.name == name) {
return true;
}
self.flags.push(RegisteredFlag { name, description });
false
}
pub fn register_event_handler(
&mut self,
tag: EventTag,
handler: EventHandlerFn,
user_data: *mut std::ffi::c_void,
) -> bool {
let idx = tag as usize;
if idx < EVENT_TAG_COUNT {
self.handlers[idx].push(RegisteredHandler { handler, user_data });
}
false
}
pub fn register_resources_discover(
&mut self,
handler: ResourcesDiscoverFn,
plugin_free_string: FreeStringFn,
user_data: *mut std::ffi::c_void,
) -> bool {
self.resources_discover.push(ResourcesDiscoverHandler {
handler,
plugin_free_string,
user_data,
});
false
}
pub fn register_provider(&mut self, provider: RegisteredProvider) -> bool {
if self
.providers
.iter()
.any(|p| p.provider_id == provider.provider_id)
{
return true;
}
self.providers.push(provider);
false
}
pub fn register_renderer(&mut self, renderer: RegisteredRenderer) -> bool {
if self
.renderers
.iter()
.any(|r| r.kind == renderer.kind && r.name == renderer.name)
{
return true;
}
self.renderers.push(renderer);
false
}
pub fn snapshot(&self) -> RegistrySnapshot {
RegistrySnapshot {
tools: self
.tools
.iter()
.map(|t| ExtensionTool {
tool: t.tool.clone(),
handle: t.handle,
})
.collect(),
commands: self.commands.clone(),
flags: self.flags.clone(),
handlers: self.handlers.clone(),
resources_discover: self.resources_discover.clone(),
providers: self.providers.clone(),
renderers: self.renderers.clone(),
active: Arc::clone(&self.active),
}
}
pub fn invalidate(&self) {
self.active.store(false, Ordering::SeqCst);
}
pub fn is_active(&self) -> bool {
self.active.load(Ordering::SeqCst)
}
pub(crate) fn absorb(&mut self, mut other: ExtensionRegistry) {
for t in other.tools.drain(..) {
if self.tools.iter().any(|x| x.tool.name == t.tool.name) {
continue;
}
self.tools.push(t);
}
for c in other.commands.drain(..) {
if self.commands.iter().any(|x| x.name == c.name) {
continue;
}
self.commands.push(c);
}
for flag in other.flags.drain(..) {
if self.flags.iter().any(|existing| existing.name == flag.name) {
continue;
}
self.flags.push(flag);
}
for (tag_idx, handlers) in other.handlers.iter_mut().enumerate() {
self.handlers[tag_idx].append(handlers);
}
self.resources_discover
.append(&mut other.resources_discover);
for p in other.providers.drain(..) {
if self
.providers
.iter()
.any(|x| x.provider_id == p.provider_id)
{
continue;
}
self.providers.push(p);
}
for r in other.renderers.drain(..) {
if self
.renderers
.iter()
.any(|x| x.kind == r.kind && x.name == r.name)
{
continue;
}
self.renderers.push(r);
}
}
}
pub struct RegistrySnapshot {
tools: Vec<ExtensionTool>,
commands: Vec<RegisteredCommand>,
flags: Vec<RegisteredFlag>,
handlers: [Vec<RegisteredHandler>; EVENT_TAG_COUNT],
resources_discover: Vec<ResourcesDiscoverHandler>,
providers: Vec<RegisteredProvider>,
renderers: Vec<RegisteredRenderer>,
active: Arc<AtomicBool>,
}
impl RegistrySnapshot {
pub fn tools(&self) -> &[ExtensionTool] {
&self.tools
}
pub fn into_tools(self) -> Vec<ExtensionTool> {
self.tools
}
pub fn commands(&self) -> &[RegisteredCommand] {
&self.commands
}
pub fn flags(&self) -> &[RegisteredFlag] {
&self.flags
}
pub fn handlers_for(&self, tag: EventTag) -> &[RegisteredHandler] {
let idx = tag as usize;
if idx < EVENT_TAG_COUNT {
&self.handlers[idx]
} else {
&[]
}
}
pub fn resources_discover(&self) -> &[ResourcesDiscoverHandler] {
&self.resources_discover
}
pub fn providers(&self) -> &[RegisteredProvider] {
&self.providers
}
pub fn renderers(&self) -> &[RegisteredRenderer] {
&self.renderers
}
pub fn renderers_of(&self, kind: RegisteredRendererKind) -> Vec<RegisteredRenderer> {
self.renderers
.iter()
.filter(|r| r.kind == kind)
.cloned()
.collect()
}
pub fn entries(&self) -> Vec<RegistryEntry> {
let mut v: Vec<RegistryEntry> = Vec::new();
for t in &self.tools {
v.push(RegistryEntry::Tool(ExtensionTool {
tool: t.tool.clone(),
handle: t.handle,
}));
}
for c in &self.commands {
v.push(RegistryEntry::Command(c.clone()));
}
for flag in &self.flags {
v.push(RegistryEntry::Flag(flag.clone()));
}
v
}
pub fn is_active(&self) -> bool {
self.active.load(Ordering::SeqCst)
}
pub fn active_flag(&self) -> &Arc<AtomicBool> {
&self.active
}
}
pub fn assert_active(active: &Arc<AtomicBool>) -> bool {
active.load(Ordering::SeqCst)
}
pub(crate) fn stale_error() -> AgentError {
AgentError::State("extensions registry is stale (session swapped/reloaded)".into())
}
#[allow(dead_code)]
fn _ensure_handler_accessor_used(snap: &RegistrySnapshot) {
let _ = snap.handlers_for(EventTag::MessageEnd);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cli_flags_snapshot_and_merge_first_wins() {
let mut first = ExtensionRegistry::new();
assert!(!first.register_flag("server".into(), "Start server".into()));
assert!(first.register_flag("server".into(), "Different description".into()));
let mut second = ExtensionRegistry::new();
assert!(!second.register_flag("server".into(), "Second server".into()));
assert!(!second.register_flag("port".into(), "Listen port".into()));
first.absorb(second);
let snapshot = first.snapshot();
let flags = snapshot.flags();
assert_eq!(flags.len(), 2);
assert_eq!(flags[0].name, "server");
assert_eq!(flags[0].description, "Start server");
assert_eq!(flags[1].name, "port");
}
}