use std::fmt;
use std::io::{self, Write};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use asupersync::Cx;
use fastmcp_core::McpRequestCancellation;
use fastmcp_protocol::http_headers::AdmittedToolHeaderSchema;
use fastmcp_protocol::protocol_policy::ProtocolEra;
use fastmcp_protocol::{
AdmittedSchema, CoreRequest, CoreResult, FinalTool, RequestId, admit_final_schema,
};
use serde_json::Value;
use super::managed::ManagedOAuthSession;
use super::rpc::tool_headers::ManagedToolHeaderError;
use super::rpc::{ManagedCoreCall, ManagedCoreError, ManagedCoreEvent, ManagedCoreLimits};
use crate::http_executor::parameter_headers::{ReviewedToolHeaders, ToolHeaderDispatchError};
pub mod catalog;
pub mod headers;
pub mod interaction;
mod validity;
use validity::await_validity;
pub const MAX_MANAGED_TOOL_SCHEMA_BYTES: usize = 512 * 1024;
pub const MAX_MANAGED_TOOL_NAME_BYTES: usize = 1024;
#[derive(Debug)]
pub enum ManagedToolError {
InvalidDefinition,
SchemaTooLarge,
InvalidInputSchema,
InvalidOutputSchema,
RequestMismatch,
InvalidArguments,
InvalidResult,
MissingStructuredOutput,
InvalidStructuredOutput,
HeaderBindingMismatch,
Invalidated,
Closed,
Headers(ToolHeaderDispatchError),
Core(ManagedCoreError),
}
impl fmt::Display for ManagedToolError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
Self::InvalidDefinition => "invalid managed tool definition",
Self::SchemaTooLarge => "managed tool schemas exceed the retained-byte limit",
Self::InvalidInputSchema => "managed tool input schema failed admission",
Self::InvalidOutputSchema => "managed tool output schema failed admission",
Self::RequestMismatch => "request does not match the bound modern tool",
Self::InvalidArguments => "tool arguments do not satisfy the admitted input schema",
Self::InvalidResult => "managed tool result failed protocol admission",
Self::MissingStructuredOutput => {
"successful tool result omitted required structured output"
}
Self::InvalidStructuredOutput => {
"tool structured output does not satisfy its admitted schema"
}
Self::HeaderBindingMismatch => "header review does not match the bound tool contract",
Self::Invalidated => "managed tool contract has been invalidated",
Self::Closed => "managed tool call is closed",
Self::Headers(error) => return fmt::Display::fmt(error, f),
Self::Core(error) => return fmt::Display::fmt(error, f),
})
}
}
impl std::error::Error for ManagedToolError {}
impl From<ManagedCoreError> for ManagedToolError {
fn from(error: ManagedCoreError) -> Self {
Self::Core(error)
}
}
impl From<ManagedToolHeaderError> for ManagedToolError {
fn from(error: ManagedToolHeaderError) -> Self {
match error {
ManagedToolHeaderError::Headers(error) => Self::Headers(error),
ManagedToolHeaderError::Core(error) => Self::Core(error),
}
}
}
struct ToolContract {
name: String,
input: AdmittedToolHeaderSchema,
output: Option<AdmittedSchema>,
invalidated: AtomicBool,
invalidation: McpRequestCancellation,
catalog_invalidated: Option<Arc<AtomicBool>>,
source_contract: Option<Arc<ToolContract>>,
}
impl ToolContract {
fn admit(tool: FinalTool) -> Result<Self, ManagedToolError> {
if tool.name.is_empty()
|| tool.name.len() > MAX_MANAGED_TOOL_NAME_BYTES
|| tool.name.bytes().any(|byte| byte < 0x20 || byte == 0x7f)
{
return Err(ManagedToolError::InvalidDefinition);
}
if tool.input_schema.get("type").and_then(Value::as_str) != Some("object") {
return Err(ManagedToolError::InvalidInputSchema);
}
if tool
.output_schema
.as_ref()
.is_some_and(|schema| !schema.is_object())
{
return Err(ManagedToolError::InvalidOutputSchema);
}
let input = AdmittedToolHeaderSchema::admit(tool.input_schema)
.map_err(|_| ManagedToolError::InvalidInputSchema)?;
let output = tool
.output_schema
.map(admit_final_schema)
.transpose()
.map_err(|_| ManagedToolError::InvalidOutputSchema)?;
let mut bytes = SchemaBytes(0);
serde_json::to_writer(&mut bytes, input.schema())
.map_err(|_| ManagedToolError::SchemaTooLarge)?;
if let Some(output) = &output {
serde_json::to_writer(&mut bytes, output.schema())
.map_err(|_| ManagedToolError::SchemaTooLarge)?;
}
Ok(Self {
name: tool.name,
input,
output,
invalidated: AtomicBool::new(false),
invalidation: McpRequestCancellation::new(),
catalog_invalidated: None,
source_contract: None,
})
}
fn invalidate(&self) {
self.invalidated.store(true, Ordering::Release);
self.invalidation.cancel();
}
fn is_invalidated(&self) -> bool {
self.invalidated.load(Ordering::Acquire)
|| self
.catalog_invalidated
.as_ref()
.is_some_and(|flag| flag.load(Ordering::Acquire))
|| self
.source_contract
.as_ref()
.is_some_and(|source| source.is_invalidated())
}
fn check(&self) -> Result<(), ManagedToolError> {
if self.is_invalidated() {
Err(ManagedToolError::Invalidated)
} else {
Ok(())
}
}
fn validate_request(&self, request: &CoreRequest) -> Result<(), ManagedToolError> {
self.check()?;
if request.era() != ProtocolEra::Modern2026 || request.method() != "tools/call" {
return Err(ManagedToolError::RequestMismatch);
}
let params = request
.encode_params()
.map_err(|_| ManagedToolError::InvalidArguments)?
.ok_or(ManagedToolError::InvalidArguments)?;
if params.get("name").and_then(Value::as_str) != Some(self.name.as_str()) {
return Err(ManagedToolError::RequestMismatch);
}
let empty = Value::Object(serde_json::Map::new());
let arguments = params.get("arguments").unwrap_or(&empty);
if !arguments.is_object() {
return Err(ManagedToolError::InvalidArguments);
}
self.input
.validate(arguments)
.map_err(|_| ManagedToolError::InvalidArguments)?;
self.check()
}
fn validate_result(&self, result: &CoreResult) -> Result<(), ManagedToolError> {
self.check()?;
let Some(output) = &self.output else {
return Ok(());
};
let encoded = result
.encode()
.map_err(|_| ManagedToolError::InvalidResult)?;
let value: Value =
serde_json::from_str(&encoded).map_err(|_| ManagedToolError::InvalidResult)?;
match value.get("resultType").and_then(Value::as_str) {
Some("input_required") => return self.check(),
Some("complete") => {}
_ => return Err(ManagedToolError::InvalidResult),
}
if value.get("isError").and_then(Value::as_bool) == Some(true) {
return self.check();
}
let structured = value
.get("structuredContent")
.ok_or(ManagedToolError::MissingStructuredOutput)?;
output
.validate(structured)
.map_err(|_| ManagedToolError::InvalidStructuredOutput)?;
self.check()
}
}
struct SchemaBytes(usize);
impl Write for SchemaBytes {
fn write(&mut self, buffer: &[u8]) -> io::Result<usize> {
if buffer.len() > MAX_MANAGED_TOOL_SCHEMA_BYTES.saturating_sub(self.0) {
return Err(io::Error::other("managed tool schema byte limit"));
}
self.0 += buffer.len();
Ok(buffer.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[derive(Clone)]
pub struct ManagedToolClient {
session: ManagedOAuthSession,
contract: Arc<ToolContract>,
header_review: Option<Arc<ReviewedToolHeaders>>,
}
impl fmt::Debug for ManagedToolClient {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ManagedToolClient")
.field("has_output_schema", &self.contract.output.is_some())
.field("parameter_headers", &self.header_review.is_some())
.field("invalidated", &self.is_invalidated())
.finish_non_exhaustive()
}
}
impl ManagedToolClient {
pub fn new(session: ManagedOAuthSession, tool: FinalTool) -> Result<Self, ManagedToolError> {
let contract = Arc::new(ToolContract::admit(tool)?);
Ok(Self {
session,
contract,
header_review: None,
})
}
pub fn tool_name(&self) -> &str {
&self.contract.name
}
pub fn invalidate(&self) {
self.contract.invalidate();
}
pub fn is_invalidated(&self) -> bool {
self.contract.is_invalidated()
}
pub fn validate_request(&self, request: &CoreRequest) -> Result<(), ManagedToolError> {
self.contract.validate_request(request)
}
pub async fn request(
&self,
cx: &Cx,
request: CoreRequest,
request_id: RequestId,
limits: ManagedCoreLimits,
) -> Result<ManagedToolCall, ManagedToolError> {
Box::pin(self.request_with_cancellation(
cx,
&McpRequestCancellation::new(),
request,
request_id,
limits,
))
.await
}
pub async fn request_with_cancellation(
&self,
cx: &Cx,
cancellation: &McpRequestCancellation,
request: CoreRequest,
request_id: RequestId,
limits: ManagedCoreLimits,
) -> Result<ManagedToolCall, ManagedToolError> {
check_tool_call(cx, cancellation, &self.contract)?;
self.contract.validate_request(&request)?;
check_tool_call(cx, cancellation, &self.contract)?;
let call = match self.header_review.as_deref() {
Some(reviewed) => {
Box::pin(await_validity(
cx,
cancellation,
&self.contract,
self.session.request_tool_with_headers_and_cancellation(
cx,
cancellation,
request,
request_id,
reviewed,
limits,
),
))
.await??
}
None => {
Box::pin(await_validity(
cx,
cancellation,
&self.contract,
self.session.request_core_with_cancellation(
cx,
cancellation,
request,
request_id,
limits,
),
))
.await??
}
};
check_tool_call(cx, cancellation, &self.contract)?;
Ok(ManagedToolCall {
call: Some(call),
contract: self.contract.clone(),
cancellation: cancellation.clone(),
finished: false,
})
}
}
pub struct ManagedToolCall {
call: Option<ManagedCoreCall>,
contract: Arc<ToolContract>,
cancellation: McpRequestCancellation,
finished: bool,
}
impl ManagedToolCall {
pub fn close(&mut self) {
self.call = None;
}
pub async fn next_event(
&mut self,
cx: &Cx,
) -> Result<Option<ManagedCoreEvent>, ManagedToolError> {
if self.finished {
return Ok(None);
}
let mut call = self.call.take().ok_or(ManagedToolError::Closed)?;
check_tool_call(cx, &self.cancellation, &self.contract)?;
let event = Box::pin(await_validity(
cx,
&self.cancellation,
&self.contract,
call.next_event(cx),
))
.await??
.ok_or(ManagedCoreError::MissingTerminal)?;
check_tool_call(cx, &self.cancellation, &self.contract)?;
match &event {
ManagedCoreEvent::Result(result) => {
self.contract.validate_result(result)?;
check_tool_call(cx, &self.cancellation, &self.contract)?;
self.finished = true;
}
ManagedCoreEvent::Notification(_) => self.call = Some(call),
}
Ok(Some(event))
}
}
fn check_tool_call(
cx: &Cx,
cancellation: &McpRequestCancellation,
contract: &ToolContract,
) -> Result<(), ManagedToolError> {
if cancellation.is_cancel_requested() || cx.checkpoint().is_err() {
return Err(ManagedCoreError::Cancelled.into());
}
contract.check()
}
#[cfg(test)]
mod tests;
#[cfg(test)]
mod header_schema_tests;