pub use mcpkit_core::tasks::{
DEFAULT_TASK_TTL_MS, RELATED_TASK_META_KEY, TaskEvent, TaskHandle, TaskManager, TaskObserver,
TaskPayload, TaskRoute, TaskState, route_task_store,
};
pub trait NotificationSink: Send + Sync {
fn publish(&self, notification: Notification);
}
impl NotificationSink for crate::server::ServerState {
fn publish(&self, notification: Notification) {
self.publish_notification(notification);
}
}
impl NotificationSink for crate::streams::StreamRegistry {
fn publish(&self, notification: Notification) {
match serde_json::to_string(&mcpkit_core::protocol::Message::Notification(notification)) {
Ok(json) => {
let _ = self.send("message", json);
}
Err(e) => tracing::warn!(error = ?e, "failed to serialize ambient notification"),
}
}
}
pub struct TaskStatusNotifier {
sink: Arc<dyn NotificationSink>,
}
impl std::fmt::Debug for TaskStatusNotifier {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TaskStatusNotifier").finish_non_exhaustive()
}
}
impl TaskStatusNotifier {
#[must_use]
pub const fn new(sink: Arc<dyn NotificationSink>) -> Self {
Self { sink }
}
}
impl TaskObserver for TaskStatusNotifier {
fn on_task_event(&self, event: &TaskEvent) {
let params = TaskStatusNotificationParams::from(event.task.clone());
match serde_json::to_value(params) {
Ok(params) => self.sink.publish(Notification::with_params(
crate::router::notifications::TASK_STATUS,
params,
)),
Err(e) => {
tracing::warn!(error = ?e, "failed to serialize task status notification");
}
}
}
}
#[must_use]
pub fn session_task_store(
streams: &Arc<crate::streams::StreamRegistry>,
default_ttl_ms: Option<u64>,
) -> Arc<TaskManager> {
let store = Arc::new(TaskManager::with_default_ttl(default_ttl_ms));
let sink: Arc<dyn NotificationSink> = Arc::<crate::streams::StreamRegistry>::clone(streams);
let _ = store.set_observer(Arc::new(TaskStatusNotifier::new(sink)));
store
}
use crate::context::Context;
use crate::handler::TaskHandler;
use mcpkit_core::error::McpError;
use mcpkit_core::protocol::Notification;
use mcpkit_core::types::task::{
CancelTaskResult, GetTaskResult, ListTasksResult, TaskId, TaskStatusNotificationParams,
};
use std::sync::Arc;
pub struct TaskService {
manager: Arc<TaskManager>,
}
impl Default for TaskService {
fn default() -> Self {
Self::new()
}
}
impl TaskService {
#[must_use]
pub fn new() -> Self {
Self {
manager: Arc::new(TaskManager::new()),
}
}
#[must_use]
pub const fn manager(&self) -> &Arc<TaskManager> {
&self.manager
}
#[must_use]
pub fn create(&self) -> TaskHandle {
self.manager.create(None)
}
}
impl TaskHandler for TaskService {
async fn list_tasks(&self, _ctx: &Context<'_>) -> Result<ListTasksResult, McpError> {
Ok(self.manager.list().into())
}
async fn get_task(
&self,
task_id: &TaskId,
_ctx: &Context<'_>,
) -> Result<Option<GetTaskResult>, McpError> {
Ok(self
.manager
.get(task_id)
.map(|s| GetTaskResult::from(s.task)))
}
async fn cancel_task(
&self,
task_id: &TaskId,
_ctx: &Context<'_>,
) -> Result<Option<CancelTaskResult>, McpError> {
if self.manager.get(task_id).is_none() {
return Ok(None);
}
self.manager.cancel(task_id)?;
Ok(self
.manager
.get(task_id)
.map(|s| CancelTaskResult::from(s.task)))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_task_service_handler() -> Result<(), Box<dyn std::error::Error>> {
let service = TaskService::new();
let handle = service.create();
let task_id = handle.id().clone();
assert_eq!(service.manager().list().len(), 1);
assert!(service.manager().get(&task_id).is_some());
Ok(())
}
}
#[cfg(test)]
mod notifier_tests {
use super::*;
use crate::streams::{StreamConfig, StreamRegistry};
#[tokio::test]
async fn session_task_store_publishes_transitions_onto_the_registry() {
let streams = Arc::new(StreamRegistry::new(StreamConfig::default()));
let (mut handle, _prime) = streams.open("message", "{}".to_string());
let store = session_task_store(&streams, None);
let task = store.create(None);
task.complete(serde_json::json!({"ok": true}))
.expect("complete");
let event = handle.recv().await.expect("an event");
let json: serde_json::Value = serde_json::from_str(&event.data).expect("json");
assert_eq!(json["method"], "notifications/tasks/status");
assert_eq!(json["params"]["status"], "completed");
assert_eq!(json["params"]["taskId"], task.id().as_str());
assert!(
json["params"]["_meta"].is_null(),
"must not carry related-task _meta: {json}"
);
}
#[tokio::test]
async fn publishing_with_no_live_stream_is_a_no_op() {
let streams = Arc::new(StreamRegistry::new(StreamConfig::default()));
let store = session_task_store(&streams, None);
let task = store.create(None);
task.complete(serde_json::json!({})).expect("complete");
assert!(!streams.has_live_stream());
}
}