use std::collections::BTreeMap;
use std::fmt;
use std::future::{Future, poll_fn};
use std::io::{self, Write};
use std::pin::Pin;
use std::task::Poll;
use std::time::Duration;
use asupersync::Cx;
use asupersync::channel::oneshot;
use asupersync::time::Sleep;
use asupersync::types::Time;
use fastmcp_core::McpRequestCancellation;
use fastmcp_protocol::common_types::SamplingContentBlock;
use fastmcp_protocol::sampling::{
SamplingToolLoop, SamplingToolLoopError, SamplingToolLoopLimits, SamplingToolLoopStep,
};
use fastmcp_protocol::{
AdmittedSchema, FinalCreateMessageResult, FinalEmbeddedCreateMessageParams, admit_final_schema,
};
pub mod inputs;
pub type SamplingHostFuture<'a, T> =
Pin<Box<dyn Future<Output = Result<T, SamplingHostError>> + Send + 'a>>;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SamplingHostError {
Denied,
Failed,
}
pub trait SamplingHost: Send {
fn sample<'a>(
&'a mut self,
cx: &'a Cx,
cancellation: &'a McpRequestCancellation,
request: &'a FinalEmbeddedCreateMessageParams,
) -> SamplingHostFuture<'a, FinalCreateMessageResult>;
fn approve_tools<'a>(
&'a mut self,
cx: &'a Cx,
cancellation: &'a McpRequestCancellation,
calls: &'a [SamplingContentBlock],
) -> SamplingHostFuture<'a, ()>;
fn execute_tool<'a>(
&'a mut self,
cx: &'a Cx,
cancellation: &'a McpRequestCancellation,
call: &'a SamplingContentBlock,
) -> SamplingHostFuture<'a, SamplingContentBlock>;
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SamplingStage {
Model,
Approval,
Tool,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SamplingRunError {
InvalidLimits,
RuntimeUnavailable,
Cancelled,
TimedOut,
Host {
stage: SamplingStage,
reason: SamplingHostError,
},
Protocol(SamplingToolLoopError),
ToolResultByteLimit,
}
impl fmt::Display for SamplingRunError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "sampling execution: {self:?}")
}
}
impl std::error::Error for SamplingRunError {}
impl From<SamplingToolLoopError> for SamplingRunError {
fn from(error: SamplingToolLoopError) -> Self {
Self::Protocol(error)
}
}
#[derive(Clone, Copy, Debug)]
pub struct SamplingRunLimits {
conversation: SamplingToolLoopLimits,
timeout: Duration,
tool_result_bytes: usize,
}
impl SamplingRunLimits {
pub fn new(
conversation: SamplingToolLoopLimits,
timeout: Duration,
tool_result_bytes: usize,
) -> Result<Self, SamplingRunError> {
if timeout.is_zero()
|| timeout > Duration::from_secs(3600)
|| tool_result_bytes == 0
|| tool_result_bytes > 16 * 1024 * 1024
{
return Err(SamplingRunError::InvalidLimits);
}
Ok(Self {
conversation,
timeout,
tool_result_bytes,
})
}
}
impl Default for SamplingRunLimits {
fn default() -> Self {
Self {
conversation: SamplingToolLoopLimits::default(),
timeout: Duration::from_secs(300),
tool_result_bytes: 4 * 1024 * 1024,
}
}
}
pub struct SamplingRunResult {
pub response: FinalCreateMessageResult,
pub model_rounds: usize,
pub executed_tools: usize,
}
pub async fn run_sampling_tool_loop<H: SamplingHost + ?Sized>(
cx: &Cx,
cancellation: &McpRequestCancellation,
request: FinalEmbeddedCreateMessageParams,
limits: SamplingRunLimits,
host: &mut H,
) -> Result<SamplingRunResult, SamplingRunError> {
let deadline = deadline(cx, cancellation, limits.timeout)?;
let mut conversation = SamplingToolLoop::new(request, limits.conversation)?;
let mut outputs: BTreeMap<String, AdmittedSchema> = BTreeMap::new();
for tool in conversation
.request()
.into_iter()
.flat_map(|request| request.tools.iter().flatten())
{
if let Some(schema) = &tool.output_schema {
outputs.insert(
tool.name.clone(),
admit_final_schema(schema.clone())
.map_err(|_| SamplingToolLoopError::InvalidSchema)?,
);
}
}
let mut executed_tools = 0;
let mut result_bytes = 0_usize;
loop {
check(cx, cancellation, deadline)?;
let request = conversation
.request()
.ok_or(SamplingToolLoopError::WrongPhase)?;
let response = within(cx, cancellation, deadline, async {
host.sample(cx, cancellation, request)
.await
.map_err(|reason| SamplingRunError::Host {
stage: SamplingStage::Model,
reason,
})
})
.await?;
match conversation.accept_response(response)? {
SamplingToolLoopStep::Complete => {
check(cx, cancellation, deadline)?;
let response = conversation
.result()
.ok_or(SamplingToolLoopError::WrongPhase)?
.clone();
check(cx, cancellation, deadline)?;
return Ok(SamplingRunResult {
response,
model_rounds: conversation.round_count(),
executed_tools,
});
}
SamplingToolLoopStep::Tools { count } => {
let calls: Vec<_> = conversation.pending_tool_calls().cloned().collect();
within(cx, cancellation, deadline, async {
host.approve_tools(cx, cancellation, &calls)
.await
.map_err(|reason| SamplingRunError::Host {
stage: SamplingStage::Approval,
reason,
})
})
.await?;
let mut results = Vec::with_capacity(count);
for call in &calls {
cooperate(cx, cancellation, deadline).await?;
let result = within(cx, cancellation, deadline, async {
host.execute_tool(cx, cancellation, call)
.await
.map_err(|reason| SamplingRunError::Host {
stage: SamplingStage::Tool,
reason,
})
})
.await?;
let remaining = limits.tool_result_bytes - result_bytes;
let bytes = encoded_size(&result, remaining)?;
validate_output(call, &result, &outputs)?;
result_bytes += bytes;
executed_tools += 1;
results.push(result);
}
conversation.submit_tool_results(results)?;
cooperate(cx, cancellation, deadline).await?;
}
}
}
}
fn validate_output(
call: &SamplingContentBlock,
result: &SamplingContentBlock,
outputs: &BTreeMap<String, AdmittedSchema>,
) -> Result<(), SamplingRunError> {
let SamplingContentBlock::ToolUse { id, name, .. } = call else {
return Err(SamplingToolLoopError::InvalidResponse.into());
};
let SamplingContentBlock::ToolResult {
tool_use_id,
structured_content,
is_error,
..
} = result
else {
return Err(SamplingToolLoopError::InvalidToolResults.into());
};
if tool_use_id != id {
return Err(SamplingToolLoopError::InvalidToolResults.into());
}
if *is_error != Some(true) {
if let Some(schema) = outputs.get(name) {
let value = structured_content
.as_ref()
.ok_or(SamplingToolLoopError::InvalidToolOutput)?;
schema
.validate(value)
.map_err(|_| SamplingToolLoopError::InvalidToolOutput)?;
}
}
Ok(())
}
fn encoded_size(value: &impl serde::Serialize, maximum: usize) -> Result<usize, SamplingRunError> {
struct Counter {
bytes: usize,
maximum: usize,
exceeded: bool,
}
impl Write for Counter {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
if bytes.len() > self.maximum - self.bytes {
self.exceeded = true;
return Err(io::Error::other("sampling result budget"));
}
self.bytes += bytes.len();
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
let mut counter = Counter {
bytes: 0,
maximum,
exceeded: false,
};
if serde_json::to_writer(&mut counter, value).is_err() {
return Err(if counter.exceeded {
SamplingRunError::ToolResultByteLimit
} else {
SamplingToolLoopError::InvalidToolResults.into()
});
}
Ok(counter.bytes)
}
fn deadline(
cx: &Cx,
cancellation: &McpRequestCancellation,
timeout: Duration,
) -> Result<Time, SamplingRunError> {
check(cx, cancellation, Time::from_nanos(u64::MAX))?;
if cx.timer_driver().is_none() {
return Err(SamplingRunError::RuntimeUnavailable);
}
let nanos = u64::try_from(timeout.as_nanos()).map_err(|_| SamplingRunError::InvalidLimits)?;
let end = cx
.now()
.as_nanos()
.checked_add(nanos)
.ok_or(SamplingRunError::InvalidLimits)?;
let deadline = Time::from_nanos(end);
check(cx, cancellation, deadline)?;
Ok(deadline)
}
fn check(
cx: &Cx,
cancellation: &McpRequestCancellation,
deadline: Time,
) -> Result<(), SamplingRunError> {
use asupersync::{CancelKind, error::ErrorKind};
if cancellation.is_cancel_requested() {
return Err(SamplingRunError::Cancelled);
}
if cx.now() >= deadline || cx.budget().deadline.is_some_and(|end| cx.now() >= end) {
return Err(SamplingRunError::TimedOut);
}
cx.checkpoint()
.map_err(|error| match cx.cancel_reason().map(|reason| reason.kind) {
Some(CancelKind::Deadline | CancelKind::Timeout) => SamplingRunError::TimedOut,
Some(_) => SamplingRunError::Cancelled,
None => match error.kind() {
ErrorKind::DeadlineExceeded | ErrorKind::CancelTimeout => {
SamplingRunError::TimedOut
}
_ => SamplingRunError::Cancelled,
},
})
}
async fn within<T>(
cx: &Cx,
cancellation: &McpRequestCancellation,
deadline: Time,
future: impl Future<Output = Result<T, SamplingRunError>>,
) -> Result<T, SamplingRunError> {
let deadline = cx
.budget()
.deadline
.map_or(deadline, |parent| parent.min(deadline));
let mut sleep = std::pin::pin!(Sleep::new(deadline));
let mut cancelled = std::pin::pin!(cancellation.cancelled());
let (_sender, mut receiver) = oneshot::channel::<()>();
let mut caller_cancelled = std::pin::pin!(receiver.recv(cx));
let mut future = std::pin::pin!(future);
poll_fn(|task| {
let _caller = Cx::set_current(Some(cx.clone()));
check(cx, cancellation, deadline)?;
if cancelled.as_mut().poll(task).is_ready()
|| caller_cancelled.as_mut().poll(task).is_ready()
{
return Poll::Ready(Err(SamplingRunError::Cancelled));
}
if sleep.as_mut().poll(task).is_ready() {
return Poll::Ready(Err(SamplingRunError::TimedOut));
}
let result = future.as_mut().poll(task);
check(cx, cancellation, deadline)?;
result
})
.await
}
async fn cooperate(
cx: &Cx,
cancellation: &McpRequestCancellation,
deadline: Time,
) -> Result<(), SamplingRunError> {
let mut yielded = false;
within(
cx,
cancellation,
deadline,
poll_fn(|task| {
if yielded {
Poll::Ready(Ok(()))
} else {
yielded = true;
task.waker().wake_by_ref();
Poll::Pending
}
}),
)
.await
}