pub mod mixed;
use std::fmt;
use asupersync::Cx;
use fastmcp_core::McpRequestCancellation;
use fastmcp_protocol::{
CoreRequest, FINAL_CLIENT_CAPABILITIES_META_KEY, FinalCreateMessageResult,
FinalEmbeddedCreateMessageParams, FinalEmbeddedInputRequest, FinalEmbeddedInputResponse,
FinalInputResponses, InputRequiredResult, RequestId, exact_json_to_serde,
};
use super::{
SamplingContentBlock, SamplingHost, SamplingHostError, SamplingHostFuture, SamplingRunError,
SamplingRunLimits, SamplingToolLoop, check, deadline, encoded_size, run_sampling_tool_loop,
within,
};
use crate::http_auth::rpc::interaction::{
ManagedInputReply, ManagedInteractionError, admit_embedded_input, validate_initial,
};
const MAX_BATCH_BYTES: usize = 16 * 1024 * 1024;
#[derive(Clone, Copy, Debug)]
pub struct SamplingInputLimits {
run: SamplingRunLimits,
inputs: usize,
model_rounds: usize,
tool_calls: usize,
input_bytes: usize,
reply_bytes: usize,
}
impl SamplingInputLimits {
pub fn new(
run: SamplingRunLimits,
inputs: usize,
model_rounds: usize,
tool_calls: usize,
input_bytes: usize,
reply_bytes: usize,
) -> Result<Self, SamplingInputError> {
if inputs > 128
|| model_rounds > 1024
|| tool_calls > 4096
|| !(2..=MAX_BATCH_BYTES).contains(&input_bytes)
|| !(2..=MAX_BATCH_BYTES).contains(&reply_bytes)
{
return Err(SamplingInputError::InvalidLimits);
}
Ok(Self {
run,
inputs,
model_rounds,
tool_calls,
input_bytes,
reply_bytes,
})
}
}
impl Default for SamplingInputLimits {
fn default() -> Self {
Self {
run: SamplingRunLimits::default(),
inputs: 8,
model_rounds: 64,
tool_calls: 256,
input_bytes: 4 * 1024 * 1024,
reply_bytes: 4 * 1024 * 1024,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SamplingInputError {
InvalidLimits,
InvalidRequestId,
InvalidRequest,
InvalidInput,
UnsupportedInput,
CapabilityNotAdvertised,
InputLimit,
InputByteLimit,
ReplyByteLimit,
ModelRoundLimit,
ToolCallLimit,
ToolResultByteLimit,
Run(SamplingRunError),
}
impl fmt::Display for SamplingInputError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "sampling input resolution: {self:?}")
}
}
impl std::error::Error for SamplingInputError {}
impl From<SamplingRunError> for SamplingInputError {
fn from(error: SamplingRunError) -> Self {
Self::Run(error)
}
}
pub async fn resolve_sampling_inputs<H: SamplingHost + ?Sized>(
cx: &Cx,
cancellation: &McpRequestCancellation,
original: &CoreRequest,
input: InputRequiredResult,
request_id: RequestId,
limits: SamplingInputLimits,
host: &mut H,
) -> Result<ManagedInputReply, SamplingInputError> {
request_id
.validate()
.map_err(|_| SamplingInputError::InvalidRequestId)?;
validate_initial(original).map_err(|_| SamplingInputError::InvalidRequest)?;
let original_params = original
.encode_params()
.map_err(|_| SamplingInputError::InvalidRequest)?
.ok_or(SamplingInputError::InvalidRequest)?;
let capabilities = &original_params["_meta"][FINAL_CLIENT_CAPABILITIES_META_KEY];
let context_advertised = capabilities["sampling"]
.get("context")
.is_some_and(serde_json::Value::is_object);
let end = deadline(cx, cancellation, limits.run.timeout)?;
let Some(map) = input.input_requests() else {
check(cx, cancellation, end)?;
return Ok(ManagedInputReply {
request_id,
input_responses: None,
});
};
if map.members().len() > limits.inputs {
return Err(SamplingInputError::InputLimit);
}
if map.members().len() > limits.model_rounds {
return Err(SamplingInputError::ModelRoundLimit);
}
let mut requests = Vec::with_capacity(map.members().len());
let mut input_bytes = 2_usize; let mut context_ignored = false;
for (index, member) in map.members().iter().enumerate() {
check(cx, cancellation, end)?;
let value =
exact_json_to_serde(&member.value).map_err(|_| SamplingInputError::InvalidInput)?;
let key_bytes = encoded_size(&member.name, limits.input_bytes - input_bytes)
.map_err(|_| SamplingInputError::InputByteLimit)?;
let value_bytes = encoded_size(&value, limits.input_bytes - input_bytes)
.map_err(|_| SamplingInputError::InputByteLimit)?;
input_bytes = input_bytes
.checked_add(key_bytes)
.and_then(|n| n.checked_add(value_bytes))
.and_then(|n| n.checked_add(1 + usize::from(index != 0)))
.filter(|n| *n <= limits.input_bytes)
.ok_or(SamplingInputError::InputByteLimit)?;
let ignore_context = !context_advertised
&& value["params"]
.get("includeContext")
.is_some_and(|context| context != "none");
let descriptor =
admit_embedded_input(capabilities, value).map_err(|error| match error {
ManagedInteractionError::CapabilityNotAdvertised => {
SamplingInputError::CapabilityNotAdvertised
}
_ => SamplingInputError::InvalidInput,
})?;
let FinalEmbeddedInputRequest::Sampling(mut request) = descriptor else {
return Err(SamplingInputError::UnsupportedInput);
};
if ignore_context {
request.include_context = None;
context_ignored = true;
}
SamplingToolLoop::new(request.clone(), limits.run.conversation)
.map_err(SamplingRunError::from)?;
requests.push((member.name.clone(), request));
}
check(cx, cancellation, end)?;
if context_ignored {
log::warn!("Ignoring sampling includeContext because sampling.context was not advertised");
}
let mut budgeted = BatchHost {
host,
models: limits.model_rounds,
tools: limits.tool_calls,
result_bytes: limits.run.tool_result_bytes,
refusal: None,
};
let outcome = within(cx, cancellation, end, async {
let mut entries = Vec::with_capacity(requests.len());
let mut reply_bytes = 2_usize;
for (index, (key, request)) in requests.into_iter().enumerate() {
let run = run_sampling_tool_loop(cx, cancellation, request, limits.run, &mut budgeted)
.await?;
let key_bytes = encoded_size(&key, limits.reply_bytes - reply_bytes)?;
let value_bytes = encoded_size(&run.response, limits.reply_bytes - reply_bytes)?;
let Some(total) = reply_bytes
.checked_add(key_bytes)
.and_then(|n| n.checked_add(value_bytes))
.and_then(|n| n.checked_add(1 + usize::from(index != 0)))
.filter(|n| *n <= limits.reply_bytes)
else {
return Err(SamplingRunError::ToolResultByteLimit);
};
reply_bytes = total;
entries.push((key, FinalEmbeddedInputResponse::Sampling(run.response)));
}
Ok(entries)
})
.await;
let entries = match outcome {
Ok(entries) => entries,
Err(error) => {
if matches!(
error,
SamplingRunError::Cancelled | SamplingRunError::TimedOut
) {
return Err(error.into());
}
if let Some(refusal) = budgeted.refusal {
return Err(refusal);
}
if error == SamplingRunError::ToolResultByteLimit {
return Err(SamplingInputError::ReplyByteLimit);
}
return Err(error.into());
}
};
let responses = FinalInputResponses::try_from_entries(entries)
.map_err(|_| SamplingInputError::InvalidInput)?;
responses
.validate_against_input_required(&input)
.map_err(|_| SamplingInputError::InvalidInput)?;
check(cx, cancellation, end)?;
Ok(ManagedInputReply {
request_id,
input_responses: Some(responses),
})
}
struct BatchHost<'a, H: ?Sized> {
host: &'a mut H,
models: usize,
tools: usize,
result_bytes: usize,
refusal: Option<SamplingInputError>,
}
impl<H: SamplingHost + ?Sized> SamplingHost for BatchHost<'_, H> {
fn sample<'a>(
&'a mut self,
cx: &'a Cx,
cancellation: &'a McpRequestCancellation,
request: &'a FinalEmbeddedCreateMessageParams,
) -> SamplingHostFuture<'a, FinalCreateMessageResult> {
if self.models == 0 {
self.refusal = Some(SamplingInputError::ModelRoundLimit);
return Box::pin(std::future::ready(Err(SamplingHostError::Failed)));
}
self.models -= 1;
self.host.sample(cx, cancellation, request)
}
fn approve_tools<'a>(
&'a mut self,
cx: &'a Cx,
cancellation: &'a McpRequestCancellation,
calls: &'a [SamplingContentBlock],
) -> SamplingHostFuture<'a, ()> {
let refusal = if self.models == 0 {
Some(SamplingInputError::ModelRoundLimit)
} else if calls.len() > self.tools {
Some(SamplingInputError::ToolCallLimit)
} else {
None
};
if let Some(refusal) = refusal {
self.refusal = Some(refusal);
return Box::pin(std::future::ready(Err(SamplingHostError::Failed)));
}
self.host.approve_tools(cx, cancellation, calls)
}
fn execute_tool<'a>(
&'a mut self,
cx: &'a Cx,
cancellation: &'a McpRequestCancellation,
call: &'a SamplingContentBlock,
) -> SamplingHostFuture<'a, SamplingContentBlock> {
if self.tools == 0 {
self.refusal = Some(SamplingInputError::ToolCallLimit);
return Box::pin(std::future::ready(Err(SamplingHostError::Failed)));
}
self.tools -= 1;
Box::pin(async move {
let result = self.host.execute_tool(cx, cancellation, call).await?;
match encoded_size(&result, self.result_bytes) {
Ok(bytes) => self.result_bytes -= bytes,
Err(SamplingRunError::ToolResultByteLimit) => {
self.refusal = Some(SamplingInputError::ToolResultByteLimit);
return Err(SamplingHostError::Failed);
}
Err(_) => return Err(SamplingHostError::Failed),
}
Ok(result)
})
}
}