use crate::{ClientOptions, ConnectionOptions};
use std::{any::Any, error::Error, sync::Arc};
#[derive(Debug, thiserror::Error)]
#[error(transparent)]
pub struct PluginError(Box<dyn Error + Send + Sync>);
impl PluginError {
pub fn new(error: impl Into<Box<dyn Error + Send + Sync>>) -> Self {
Self(error.into())
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, derive_more::Display)]
#[non_exhaustive]
pub enum PluginTarget {
#[display("connection options")]
Connection,
#[display("client options")]
Client,
#[display("worker options")]
Worker,
#[display("workflow replayer options")]
WorkflowReplayer,
}
#[derive(Debug, thiserror::Error)]
#[error("plugin '{plugin_name}' failed to configure {target}: {source}")]
#[non_exhaustive]
pub struct PluginApplyError {
pub plugin_name: String,
pub target: PluginTarget,
#[source]
pub source: PluginError,
}
impl PluginApplyError {
pub fn new(plugin_name: impl Into<String>, target: PluginTarget, source: PluginError) -> Self {
Self {
plugin_name: plugin_name.into(),
target,
source,
}
}
}
pub trait ClientPlugin: Send + Sync + 'static {
fn name(&self) -> &str;
fn configure_connection_options(
&self,
_options: &mut ConnectionOptions,
) -> Result<(), PluginError> {
Ok(())
}
fn configure_client_options(&self, _options: &mut ClientOptions) -> Result<(), PluginError> {
Ok(())
}
}
pub trait WorkerPluginData: Any + Send + Sync + 'static {}
#[derive(Clone)]
pub struct ErasedClientPlugin {
client: Arc<dyn ClientPlugin>,
worker_plugins: Vec<Arc<dyn WorkerPluginData>>,
}
impl ErasedClientPlugin {
pub fn new<P: ClientPlugin>(plugin: P) -> Self {
Self {
client: Arc::new(plugin),
worker_plugins: Vec::new(),
}
}
pub fn with_worker_plugin<T: WorkerPluginData>(mut self, plugin: T) -> Self {
self.worker_plugins.push(Arc::new(plugin));
self
}
pub fn worker_plugins(&self) -> impl Iterator<Item = &dyn WorkerPluginData> {
self.worker_plugins.iter().map(AsRef::as_ref)
}
pub fn name(&self) -> &str {
self.client.name()
}
pub(crate) fn plugin(&self) -> &dyn ClientPlugin {
self.client.as_ref()
}
}
pub(crate) fn apply_connection_plugins(
client_options: &ClientOptions,
connection_options: &mut ConnectionOptions,
) -> Result<(), PluginApplyError> {
for registration in client_options.plugins() {
registration
.plugin()
.configure_connection_options(connection_options)
.map_err(|source| {
PluginApplyError::new(
registration.plugin().name(),
PluginTarget::Connection,
source,
)
})?;
}
Ok(())
}
pub(crate) fn apply_client_plugins(options: &mut ClientOptions) -> Result<(), PluginApplyError> {
if options.client_plugins_applied() {
return Ok(());
}
let plugins = options.plugins().to_vec();
for registration in plugins {
registration
.plugin()
.configure_client_options(options)
.map_err(|source| {
PluginApplyError::new(registration.plugin().name(), PluginTarget::Client, source)
})?;
}
options.mark_client_plugins_applied();
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use url::Url;
struct CountingPlugin {
connection_calls: Arc<AtomicUsize>,
client_calls: Arc<AtomicUsize>,
}
impl ClientPlugin for CountingPlugin {
fn name(&self) -> &str {
"counting"
}
fn configure_connection_options(
&self,
options: &mut ConnectionOptions,
) -> Result<(), PluginError> {
self.connection_calls.fetch_add(1, Ordering::Relaxed);
options.identity.push_str("-configured");
Ok(())
}
fn configure_client_options(&self, options: &mut ClientOptions) -> Result<(), PluginError> {
self.client_calls.fetch_add(1, Ordering::Relaxed);
options.namespace.push_str("-configured");
Ok(())
}
}
#[test]
fn plugins_follow_target_lifecycles() {
let connection_calls = Arc::new(AtomicUsize::new(0));
let client_calls = Arc::new(AtomicUsize::new(0));
let mut client_options = ClientOptions::new("namespace")
.client_plugin(CountingPlugin {
connection_calls: connection_calls.clone(),
client_calls: client_calls.clone(),
})
.build();
let mut first_connection_options =
ConnectionOptions::new(Url::parse("http://localhost:7233").unwrap())
.identity("first")
.build();
let mut second_connection_options =
ConnectionOptions::new(Url::parse("http://localhost:7233").unwrap())
.identity("second")
.build();
apply_connection_plugins(&client_options, &mut first_connection_options).unwrap();
apply_client_plugins(&mut client_options).unwrap();
apply_connection_plugins(&client_options, &mut second_connection_options).unwrap();
apply_client_plugins(&mut client_options).unwrap();
assert_eq!(connection_calls.load(Ordering::Relaxed), 2);
assert_eq!(client_calls.load(Ordering::Relaxed), 1);
assert_eq!(first_connection_options.identity, "first-configured");
assert_eq!(second_connection_options.identity, "second-configured");
assert_eq!(client_options.namespace, "namespace-configured");
}
}