use std::collections::HashMap;
use std::sync::Arc;
use parking_lot::RwLock;
#[derive(Debug, Clone)]
pub struct PluginMetadata {
pub name: String,
pub version: String,
pub description: String,
pub author: String,
}
impl PluginMetadata {
pub fn new(
name: impl Into<String>,
version: impl Into<String>,
description: impl Into<String>,
) -> Self {
Self {
name: name.into(),
version: version.into(),
description: description.into(),
author: String::new(),
}
}
pub fn with_author(mut self, author: impl Into<String>) -> Self {
self.author = author.into();
self
}
}
pub trait AiExtension: Send + Sync {
fn name(&self) -> &str;
fn execute(&self, input: &str) -> Result<String, PluginError>;
}
pub trait DialectExtension: Send + Sync {
fn dialect_name(&self) -> &str;
fn translate(&self, sql: &str) -> Result<String, PluginError>;
}
pub trait MiddlewareExtension: Send + Sync {
fn name(&self) -> &str;
fn before_query(&self, sql: &str) -> Result<String, PluginError>;
fn after_query(&self, sql: &str, result: &str) -> Result<String, PluginError>;
}
#[derive(Debug, Clone)]
pub enum PluginError {
NotFound(String),
ExecutionFailed(String),
RegistrationFailed(String),
}
impl std::fmt::Display for PluginError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PluginError::NotFound(msg) => write!(f, "Plugin not found: {}", msg),
PluginError::ExecutionFailed(msg) => write!(f, "Execution failed: {}", msg),
PluginError::RegistrationFailed(msg) => write!(f, "Registration failed: {}", msg),
}
}
}
impl std::error::Error for PluginError {}
pub trait SzOrmPlugin: Send + Sync {
fn metadata(&self) -> &PluginMetadata;
fn init(&self) -> Result<(), PluginError> {
Ok(())
}
fn ai_extension(&self) -> Option<&dyn AiExtension> {
None
}
fn dialect_extension(&self) -> Option<&dyn DialectExtension> {
None
}
fn middleware_extension(&self) -> Option<&dyn MiddlewareExtension> {
None
}
}
pub struct PluginRegistry {
plugins: RwLock<HashMap<String, Arc<dyn SzOrmPlugin>>>,
}
impl Default for PluginRegistry {
fn default() -> Self {
Self::new()
}
}
impl PluginRegistry {
pub fn new() -> Self {
Self {
plugins: RwLock::new(HashMap::new()),
}
}
pub fn register(&self, plugin: Arc<dyn SzOrmPlugin>) -> Result<(), PluginError> {
let metadata = plugin.metadata();
let name = metadata.name.clone();
plugin.init()?;
let mut plugins = self.plugins.write();
if plugins.contains_key(&name) {
return Err(PluginError::RegistrationFailed(format!(
"插件 {} 已存在",
name
)));
}
plugins.insert(name, plugin);
Ok(())
}
pub fn unregister(&self, name: &str) -> Result<(), PluginError> {
let mut plugins = self.plugins.write();
plugins
.remove(name)
.ok_or_else(|| PluginError::NotFound(name.to_string()))?;
Ok(())
}
pub fn get(&self, name: &str) -> Option<Arc<dyn SzOrmPlugin>> {
self.plugins.read().get(name).cloned()
}
pub fn list(&self) -> Vec<String> {
self.plugins.read().keys().cloned().collect()
}
pub fn len(&self) -> usize {
self.plugins.read().len()
}
pub fn is_empty(&self) -> bool {
self.plugins.read().is_empty()
}
pub fn execute_ai(&self, plugin_name: &str, input: &str) -> Result<String, PluginError> {
let plugin = self
.get(plugin_name)
.ok_or_else(|| PluginError::NotFound(plugin_name.to_string()))?;
let ext = plugin
.ai_extension()
.ok_or_else(|| PluginError::ExecutionFailed("插件无 AI 扩展".to_string()))?;
ext.execute(input)
}
pub fn translate_dialect(&self, plugin_name: &str, sql: &str) -> Result<String, PluginError> {
let plugin = self
.get(plugin_name)
.ok_or_else(|| PluginError::NotFound(plugin_name.to_string()))?;
let ext = plugin
.dialect_extension()
.ok_or_else(|| PluginError::ExecutionFailed("插件无方言扩展".to_string()))?;
ext.translate(sql)
}
pub fn before_query(&self, plugin_name: &str, sql: &str) -> Result<String, PluginError> {
let plugin = self
.get(plugin_name)
.ok_or_else(|| PluginError::NotFound(plugin_name.to_string()))?;
let ext = plugin
.middleware_extension()
.ok_or_else(|| PluginError::ExecutionFailed("插件无中间件扩展".to_string()))?;
ext.before_query(sql)
}
pub fn after_query(
&self,
plugin_name: &str,
sql: &str,
result: &str,
) -> Result<String, PluginError> {
let plugin = self
.get(plugin_name)
.ok_or_else(|| PluginError::NotFound(plugin_name.to_string()))?;
let ext = plugin
.middleware_extension()
.ok_or_else(|| PluginError::ExecutionFailed("插件无中间件扩展".to_string()))?;
ext.after_query(sql, result)
}
}