use std::future::Future;
use asupersync::Cx;
use fastmcp_core::McpRequestCancellation;
use fastmcp_protocol::{CoreResult, InputRequiredResult, ServerNotification};
use super::super::{ManagedToolError, ToolContract, await_validity, check_tool_call};
use super::{ManagedInteractionError, ManagedToolInteraction, ManagedToolInteractionError};
use crate::http_auth::rpc::interaction::ManagedInputReply;
impl ManagedToolInteraction {
pub async fn drive<R, F, N>(
self,
cx: &Cx,
resolve: R,
notify: N,
) -> Result<Box<CoreResult>, ManagedToolInteractionError>
where
R: FnMut(Box<InputRequiredResult>) -> F,
F: Future<Output = Result<ManagedInputReply, ManagedInteractionError>>,
N: FnMut(Box<ServerNotification>) -> Result<(), ManagedInteractionError>,
{
Box::pin(self.drive_selected(cx, resolve, notify, false)).await
}
pub async fn drive_partial<R, F, N>(
self,
cx: &Cx,
resolve: R,
notify: N,
) -> Result<Box<CoreResult>, ManagedToolInteractionError>
where
R: FnMut(Box<InputRequiredResult>) -> F,
F: Future<Output = Result<ManagedInputReply, ManagedInteractionError>>,
N: FnMut(Box<ServerNotification>) -> Result<(), ManagedInteractionError>,
{
Box::pin(self.drive_selected(cx, resolve, notify, true)).await
}
async fn drive_selected<R, F, N>(
mut self,
cx: &Cx,
mut resolve: R,
mut notify: N,
partial: bool,
) -> Result<Box<CoreResult>, ManagedToolInteractionError>
where
R: FnMut(Box<InputRequiredResult>) -> F,
F: Future<Output = Result<ManagedInputReply, ManagedInteractionError>>,
N: FnMut(Box<ServerNotification>) -> Result<(), ManagedInteractionError>,
{
check_tool_call(cx, &self.cancellation, &self.contract)?;
if self.finished {
return Err(ManagedToolError::Closed.into());
}
let operation = self.operation.take().ok_or(ManagedToolError::Closed)?;
let contract = self.contract.as_ref();
let cancellation = &self.cancellation;
let guarded_resolve = |input| {
let future = begin_resolution(cx, cancellation, contract, &mut resolve, input);
finish_resolution(cx, cancellation, contract, future)
};
let guarded_notify = |notification| {
deliver_notification(cx, cancellation, contract, &mut notify, notification)
};
let result = Box::pin(await_validity(cx, cancellation, contract, async {
if partial {
operation
.drive_partial(cx, guarded_resolve, guarded_notify)
.await
} else {
operation.drive(cx, guarded_resolve, guarded_notify).await
}
}))
.await?;
check_tool_call(cx, cancellation, contract)?;
let result = result?;
contract.validate_result(&result)?;
check_tool_call(cx, cancellation, contract)?;
Ok(result)
}
}
fn begin_resolution<R, F>(
cx: &Cx,
cancellation: &McpRequestCancellation,
contract: &ToolContract,
resolve: &mut R,
input: Box<InputRequiredResult>,
) -> Result<F, ManagedToolError>
where
R: FnMut(Box<InputRequiredResult>) -> F,
{
check_tool_call(cx, cancellation, contract)?;
Ok(resolve(input))
}
async fn finish_resolution<F>(
cx: &Cx,
cancellation: &McpRequestCancellation,
contract: &ToolContract,
future: Result<F, ManagedToolError>,
) -> Result<ManagedInputReply, ManagedInteractionError>
where
F: Future<Output = Result<ManagedInputReply, ManagedInteractionError>>,
{
let future = future.map_err(|_| ManagedInteractionError::AbortedByHost)?;
await_validity(cx, cancellation, contract, future)
.await
.map_err(|_| ManagedInteractionError::AbortedByHost)?
}
fn deliver_notification<N>(
cx: &Cx,
cancellation: &McpRequestCancellation,
contract: &ToolContract,
notify: &mut N,
notification: Box<ServerNotification>,
) -> Result<(), ManagedInteractionError>
where
N: FnMut(Box<ServerNotification>) -> Result<(), ManagedInteractionError>,
{
check_tool_call(cx, cancellation, contract)
.map_err(|_| ManagedInteractionError::AbortedByHost)?;
let outcome = notify(notification);
check_tool_call(cx, cancellation, contract)
.map_err(|_| ManagedInteractionError::AbortedByHost)?;
outcome
}
#[cfg(test)]
mod tests;