use anyhow::{anyhow, Result};
use mcp_protocol::types::sampling::{CreateMessageParams, CreateMessageResult};
use std::sync::Arc;
use tokio::sync::Mutex;
pub type CreateMessageCallback = Box<dyn Fn(&CreateMessageParams) -> Result<CreateMessageResult> + Send + Sync>;
pub struct SamplingManager {
create_message_callback: Arc<Mutex<Option<CreateMessageCallback>>>,
}
impl SamplingManager {
pub fn new() -> Self {
Self {
create_message_callback: Arc::new(Mutex::new(None)),
}
}
pub fn register_create_message_callback(&self, callback: CreateMessageCallback) {
let mut cb = self.create_message_callback.blocking_lock();
*cb = Some(callback);
}
pub async fn create_message(&self, params: &CreateMessageParams) -> Result<CreateMessageResult> {
let cb = self.create_message_callback.lock().await;
if cb.is_none() {
return Err(anyhow!("No create message callback registered"));
}
let callback_ref = cb.as_ref().unwrap();
callback_ref(params)
}
}
impl Default for SamplingManager {
fn default() -> Self {
Self::new()
}
}