pub mod creation;
pub mod driver;
pub mod submission;
pub mod subscriptions;
use std::fmt;
use std::io::{self, Write};
use std::time::Duration;
use asupersync::Cx;
use asupersync::types::Time;
use fastmcp_core::McpRequestCancellation;
use fastmcp_protocol::common_types::ExactNonNegativeJsonNumber;
use fastmcp_protocol::tasks_extension::{
CancelTaskParams, GetTaskParams, GetTaskResult, TASK_CANCEL, TASK_GET, TASK_UPDATE,
TASKS_EXTENSION, Task, TaskId, TaskInputLedger, TaskMethodRequest, TaskRequestMeta,
UpdateTaskParams,
};
use fastmcp_protocol::{
CoreRequest, CoreResult, ExtensionDirection, FINAL_CLIENT_CAPABILITIES_META_KEY,
FINAL_PROTOCOL_VERSION, FinalCoreResult, FinalRequestMeta, JsonRpcMessage, JsonRpcResponse,
ProgressMarker, RequestId, ServerDiscoverResult, ServerNotification,
decode_strict_jsonrpc_message, decode_strict_jsonrpc_response,
};
use serde::Deserialize;
use serde_json::{Value, json};
use super::{
CLIENT_CREDENTIALS_EXTENSION, ClientCredentialsClient, ClientCredentialsError,
ClientCredentialsSnapshot, active, admit_resource, authorize, discovery_deadline,
};
use crate::http_auth::managed::tasks::validate_task_input_shapes;
pub use crate::http_auth::managed::tasks::{
ManagedTaskEvent, ManagedTaskRequest, ManagedTasksError,
};
use crate::http_executor::{
ModernHttpExecutorError, ModernHttpRequest, ModernHttpResponseKind, ModernHttpResponseStream,
ModernHttpSseResponseStream,
};
use crate::sse::SseLimits;
#[derive(Clone, Copy, Debug)]
pub struct ClientCredentialsTasksLimits {
request_bytes: usize,
frame_bytes: usize,
records: usize,
timeout: Duration,
}
impl Default for ClientCredentialsTasksLimits {
fn default() -> Self {
Self {
request_bytes: 64 * 1024,
frame_bytes: 64 * 1024,
records: 64,
timeout: Duration::from_secs(120),
}
}
}
impl ClientCredentialsTasksLimits {
pub fn new(
request_bytes: usize,
frame_bytes: usize,
records: usize,
timeout: Duration,
) -> Result<Self, ClientCredentialsTasksError> {
if !(1..=1024 * 1024).contains(&request_bytes)
|| !(1..=1024 * 1024).contains(&frame_bytes)
|| !(1..=1024).contains(&records)
|| timeout.is_zero()
|| timeout > Duration::from_mins(15)
{
return Err(ManagedTasksError::InvalidLimits.into());
}
Ok(Self {
request_bytes,
frame_bytes,
records,
timeout,
})
}
}
#[derive(Debug)]
pub enum ClientCredentialsTasksError {
Authentication(ClientCredentialsError),
Protocol(ManagedTasksError),
}
impl fmt::Display for ClientCredentialsTasksError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Authentication(error) => fmt::Display::fmt(error, f),
Self::Protocol(error) => fmt::Display::fmt(error, f),
}
}
}
impl std::error::Error for ClientCredentialsTasksError {}
impl From<ClientCredentialsError> for ClientCredentialsTasksError {
fn from(error: ClientCredentialsError) -> Self {
match error {
ClientCredentialsError::Redirect { status } => {
Self::Protocol(ManagedTasksError::HttpStatus { status })
}
other => Self::Authentication(other),
}
}
}
impl From<ManagedTasksError> for ClientCredentialsTasksError {
fn from(error: ManagedTasksError) -> Self {
Self::Protocol(error)
}
}
#[derive(Clone)]
pub struct ClientCredentialsTasksClient {
client: ClientCredentialsClient,
metadata: Value,
limits: ClientCredentialsTasksLimits,
}
impl ClientCredentialsTasksClient {
pub fn new(
client: ClientCredentialsClient,
metadata: FinalRequestMeta,
limits: ClientCredentialsTasksLimits,
) -> Result<Self, ClientCredentialsTasksError> {
Ok(Self {
client,
metadata: task_metadata(metadata)?,
limits,
})
}
pub async fn request(
&self,
cx: &Cx,
discovery_id: RequestId,
request_id: RequestId,
request: ManagedTaskRequest,
) -> Result<ClientCredentialsTaskCall, ClientCredentialsTasksError> {
Box::pin(self.request_with_cancellation(
cx,
&McpRequestCancellation::new(),
discovery_id,
request_id,
request,
))
.await
}
pub async fn request_with_cancellation(
&self,
cx: &Cx,
cancellation: &McpRequestCancellation,
discovery_id: RequestId,
request_id: RequestId,
request: ManagedTaskRequest,
) -> Result<ClientCredentialsTaskCall, ClientCredentialsTasksError> {
let round = self.prepare_round(discovery_id, request_id, request)?;
let deadline = discovery_deadline(cx, self.limits.timeout.min(self.client.inner.timeout))
.map_err(ClientCredentialsError::from)?;
Box::pin(self.execute_round(cx, cancellation, round, deadline)).await
}
fn prepare_round(
&self,
discovery_id: RequestId,
request_id: RequestId,
request: ManagedTaskRequest,
) -> Result<PreparedRound, ClientCredentialsTasksError> {
discovery_id
.validate()
.map_err(|_| ManagedTasksError::InvalidRequest)?;
request_id
.validate()
.map_err(|_| ManagedTasksError::InvalidRequest)?;
if discovery_id.correlates_with(&request_id) {
return Err(ManagedTasksError::InvalidRequest.into());
}
let prepared = prepare(
self.client.resource().as_str(),
&self.metadata,
&request_id,
request,
self.limits,
)?;
let discovery = CoreRequest::decode(
fastmcp_protocol::protocol_policy::ProtocolEra::Modern2026,
"server/discover",
Some(&json!({"_meta": self.metadata})),
)
.map_err(|_| ManagedTasksError::InvalidRequest)?;
let params = discovery
.encode_params()
.map_err(|_| ManagedTasksError::InvalidRequest)?
.ok_or(ManagedTasksError::InvalidRequest)?;
let discovery_wire = encode(
self.client.resource().as_str(),
"server/discover",
&discovery_id,
params,
None,
self.limits.request_bytes,
)?;
Ok(PreparedRound {
discovery_id,
request_id,
prepared,
discovery,
discovery_wire,
})
}
async fn execute_round(
&self,
cx: &Cx,
cancellation: &McpRequestCancellation,
round: PreparedRound,
deadline: Time,
) -> Result<ClientCredentialsTaskCall, ClientCredentialsTasksError> {
let PreparedRound {
discovery_id,
request_id,
prepared,
discovery,
discovery_wire,
} = round;
let owner = &self.client.inner.closed;
active(cx, deadline, owner, cancellation, None, async {
Ok(Box::pin(async {
let snapshot = self
.client
.credential_with_cancellation(cx, cancellation)
.await?;
let executor = self.client.resource_http_executor();
let discovery_wire = authorize(&snapshot, discovery_wire)?;
let response = active(cx, deadline, owner, cancellation, Some(&snapshot), async {
executor
.execute_with_cancellation(cx, cancellation, &discovery_wire)
.await
.map_err(|_| ClientCredentialsError::Transport)
})
.await?;
require_success(&response)?;
if response.metadata().kind() != ModernHttpResponseKind::Json {
return Err(ManagedTasksError::Negotiation.into());
}
let bytes = active(cx, deadline, owner, cancellation, Some(&snapshot), async {
response
.read_to_end_with_cancellation(cx, cancellation, self.limits.frame_bytes)
.await
.map_err(|_| ClientCredentialsError::UnexpectedResponse)
})
.await?;
admit_composition(
&discovery,
&discovery_id,
&bytes,
&prepared.decoder,
self.limits.frame_bytes,
)?;
let wire = authorize(&snapshot, prepared.wire)?;
let response = active(cx, deadline, owner, cancellation, Some(&snapshot), async {
executor
.execute_with_cancellation(cx, cancellation, &wire)
.await
.map_err(|_| ClientCredentialsError::Transport)
})
.await?;
require_success(&response)?;
ClientCredentialsTaskCall::new(
response,
prepared.decoder,
prepared.progress,
snapshot,
owner.clone(),
cancellation.clone(),
request_id,
deadline,
self.limits,
)
})
.await)
})
.await?
}
}
fn require_success(response: &ModernHttpResponseStream) -> Result<(), ManagedTasksError> {
if response.metadata().status() == 200 {
Ok(())
} else {
Err(ManagedTasksError::HttpStatus {
status: response.metadata().status(),
})
}
}
enum Decoder {
Tool(Box<CoreRequest>),
Get(TaskId),
Update,
Cancel,
}
struct Prepared {
wire: ModernHttpRequest,
decoder: Decoder,
progress: Option<ProgressMarker>,
}
struct PreparedRound {
discovery_id: RequestId,
request_id: RequestId,
prepared: Prepared,
discovery: CoreRequest,
discovery_wire: ModernHttpRequest,
}
enum Body {
Json(ModernHttpResponseStream),
Sse(ModernHttpSseResponseStream),
}
pub struct ClientCredentialsTaskCall {
body: Option<Box<Body>>,
decoder: Decoder,
snapshot: ClientCredentialsSnapshot,
owner: McpRequestCancellation,
cancellation: McpRequestCancellation,
request_id: RequestId,
progress: Option<ProgressMarker>,
last_progress: Option<ExactNonNegativeJsonNumber>,
deadline: Time,
limits: ClientCredentialsTasksLimits,
records: usize,
finished: bool,
}
impl ClientCredentialsTaskCall {
#[allow(clippy::too_many_arguments)]
fn new(
response: ModernHttpResponseStream,
decoder: Decoder,
progress: Option<ProgressMarker>,
snapshot: ClientCredentialsSnapshot,
owner: McpRequestCancellation,
cancellation: McpRequestCancellation,
request_id: RequestId,
deadline: Time,
limits: ClientCredentialsTasksLimits,
) -> Result<Self, ClientCredentialsTasksError> {
let body = match response.metadata().kind() {
ModernHttpResponseKind::Json => Body::Json(response),
ModernHttpResponseKind::Sse => {
let framing = SseLimits::new(limits.frame_bytes, limits.frame_bytes, 64)
.ok_or(ManagedTasksError::InvalidLimits)?;
Body::Sse(
response
.into_sse_stream(framing)
.map_err(|_| ManagedTasksError::InvalidResponse)?,
)
}
_ => return Err(ManagedTasksError::InvalidResponse.into()),
};
Ok(Self {
body: Some(Box::new(body)),
decoder,
snapshot,
owner,
cancellation,
request_id,
progress,
last_progress: None,
deadline,
limits,
records: 0,
finished: false,
})
}
pub fn request_id(&self) -> &RequestId {
&self.request_id
}
pub fn credential_generation(&self) -> u64 {
self.generation()
}
fn generation(&self) -> u64 {
self.snapshot.generation()
}
pub fn close(&mut self) {
self.body = None;
}
pub async fn next_event(
&mut self,
cx: &Cx,
) -> Result<Option<ManagedTaskEvent>, ClientCredentialsTasksError> {
if self.finished {
return Ok(None);
}
let body = self.body.take().ok_or(ManagedTasksError::Closed)?;
let mut progress = self.last_progress.clone();
let (event, remaining) = active(
cx,
self.deadline,
&self.owner,
&self.cancellation,
Some(&self.snapshot),
async {
Ok(async {
if self.records >= self.limits.records {
return Err(ManagedTasksError::RecordLimit);
}
match *body {
Body::Json(response) => {
let bytes = response
.read_to_end_with_cancellation(
cx,
&self.cancellation,
self.limits.frame_bytes,
)
.await
.map_err(|error| json_body_error(&self.decoder, error))?;
Ok((
decode_result(
&self.decoder,
&bytes,
&self.request_id,
self.limits.frame_bytes,
)?,
None,
))
}
Body::Sse(mut stream) => {
let frame = stream
.next_event(cx)
.await
.map_err(|error| json_body_error(&self.decoder, error))?
.ok_or(ManagedTasksError::MissingTerminal)?;
let event = decode_record(
&self.decoder,
frame.as_bytes(),
&self.request_id,
self.limits.frame_bytes,
self.progress.as_ref(),
&mut progress,
)?;
if !matches!(&event, ManagedTaskEvent::Notification(_)) {
super::rpc::require_finite_sse_eof(cx, &mut stream)
.await
.map_err(|_| ManagedTasksError::InvalidResponse)?;
return Ok((event, None));
}
Ok((event, Some(Box::new(Body::Sse(stream)))))
}
}
}
.await)
},
)
.await??;
self.last_progress = progress;
self.records += 1;
if matches!(&event, ManagedTaskEvent::Notification(_)) {
self.body = remaining;
} else {
self.finished = true;
}
Ok(Some(event))
}
}
fn json_body_error(decoder: &Decoder, error: ModernHttpExecutorError) -> ManagedTasksError {
if matches!(decoder, Decoder::Get(_))
&& matches!(error, ModernHttpExecutorError::ResponseBodyReadFailed)
{
ManagedTasksError::MissingTerminal
} else {
ManagedTasksError::InvalidResponse
}
}
fn task_metadata(metadata: FinalRequestMeta) -> Result<Value, ManagedTasksError> {
let mut meta = serde_json::to_value(metadata).map_err(|_| ManagedTasksError::InvalidRequest)?;
let caps = meta
.get_mut(FINAL_CLIENT_CAPABILITIES_META_KEY)
.and_then(Value::as_object_mut)
.ok_or(ManagedTasksError::InvalidRequest)?;
if let Some(extensions) = caps.get("extensions") {
let extensions = extensions
.as_object()
.ok_or(ManagedTasksError::Negotiation)?;
if extensions.iter().any(|(name, value)| {
(name != TASKS_EXTENSION && name != CLIENT_CREDENTIALS_EXTENSION)
|| !value.as_object().is_some_and(serde_json::Map::is_empty)
}) {
return Err(ManagedTasksError::Negotiation);
}
}
caps.insert(
"extensions".to_owned(),
json!({TASKS_EXTENSION:{}, CLIENT_CREDENTIALS_EXTENSION:{}}),
);
CoreRequest::decode(
fastmcp_protocol::protocol_policy::ProtocolEra::Modern2026,
"server/discover",
Some(&json!({"_meta":meta})),
)
.map_err(|_| ManagedTasksError::InvalidRequest)?;
Ok(meta)
}
fn prepare(
target: &str,
metadata: &Value,
id: &RequestId,
request: ManagedTaskRequest,
limits: ClientCredentialsTasksLimits,
) -> Result<Prepared, ManagedTasksError> {
let request_meta = TaskRequestMeta {
meta: serde_json::from_value(metadata.clone())
.map_err(|_| ManagedTasksError::InvalidRequest)?,
};
let progress = metadata
.get("progressToken")
.map(|value| {
serde_json::from_value(value.clone()).map_err(|_| ManagedTasksError::InvalidRequest)
})
.transpose()?;
let (method, params, name, decoder) = match request {
ManagedTaskRequest::CallTool { name, arguments } => {
let mut params = json!({"_meta":metadata, "name":name});
if let Some(arguments) = arguments {
params["arguments"] = arguments;
}
let request = CoreRequest::decode(
fastmcp_protocol::protocol_policy::ProtocolEra::Modern2026,
"tools/call",
Some(¶ms),
)
.map_err(|_| ManagedTasksError::InvalidRequest)?;
(
"tools/call",
params,
Some(name),
Decoder::Tool(Box::new(request)),
)
}
ManagedTaskRequest::Get(id_to_get) => {
let request = TaskMethodRequest::new(
id.clone(),
TASK_GET,
GetTaskParams {
request: request_meta,
task_id: id_to_get.clone(),
},
);
let request = TaskMethodRequest::decode(
serde_json::to_value(request).map_err(|_| ManagedTasksError::InvalidRequest)?,
)
.map_err(|_| ManagedTasksError::InvalidRequest)?;
(
TASK_GET,
serde_json::to_value(request.params)
.map_err(|_| ManagedTasksError::InvalidRequest)?,
None,
Decoder::Get(id_to_get),
)
}
ManagedTaskRequest::Update {
task,
input_responses,
} => {
let Task::InputRequired {
base,
input_requests,
} = *task
else {
return Err(ManagedTasksError::UpdateRequiresInput);
};
let ledger = TaskInputLedger::from_requests(&input_requests)
.map_err(|_| ManagedTasksError::InvalidInputResponses)?;
ledger
.validate_responses(&input_responses)
.map_err(|_| ManagedTasksError::InvalidInputResponses)?;
let request = TaskMethodRequest::new(
id.clone(),
TASK_UPDATE,
UpdateTaskParams {
request: request_meta,
task_id: base.task_id,
input_responses,
},
);
let request = TaskMethodRequest::decode_update(
serde_json::to_value(request).map_err(|_| ManagedTasksError::InvalidRequest)?,
&ledger,
)
.map_err(|_| ManagedTasksError::InvalidInputResponses)?;
(
TASK_UPDATE,
serde_json::to_value(request.params)
.map_err(|_| ManagedTasksError::InvalidRequest)?,
None,
Decoder::Update,
)
}
ManagedTaskRequest::Cancel(task_id) => {
let request = TaskMethodRequest::new(
id.clone(),
TASK_CANCEL,
CancelTaskParams {
request: request_meta,
task_id,
},
);
let request = TaskMethodRequest::decode_cancel(
serde_json::to_value(request).map_err(|_| ManagedTasksError::InvalidRequest)?,
)
.map_err(|_| ManagedTasksError::InvalidRequest)?;
(
TASK_CANCEL,
serde_json::to_value(request.params)
.map_err(|_| ManagedTasksError::InvalidRequest)?,
None,
Decoder::Cancel,
)
}
};
Ok(Prepared {
wire: encode(target, method, id, params, name, limits.request_bytes)?,
decoder,
progress,
})
}
fn encode(
target: &str,
method: &str,
id: &RequestId,
params: Value,
name: Option<String>,
maximum: usize,
) -> Result<ModernHttpRequest, ManagedTasksError> {
id.validate()
.map_err(|_| ManagedTasksError::InvalidRequest)?;
let mut body = BoundedBody {
bytes: Vec::new(),
maximum,
};
serde_json::to_writer(
&mut body,
&json!({"jsonrpc":"2.0", "id":id, "method":method, "params":params}),
)
.map_err(|_| ManagedTasksError::RequestTooLarge)?;
ModernHttpRequest::new(target, body.bytes, FINAL_PROTOCOL_VERSION, method, name)
.map_err(|_| ManagedTasksError::InvalidRequest)
}
struct BoundedBody {
bytes: Vec<u8>,
maximum: usize,
}
impl Write for BoundedBody {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
if bytes.len() > self.maximum.saturating_sub(self.bytes.len()) {
return Err(io::Error::other("machine Tasks request limit"));
}
self.bytes.extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
fn response_source(
bytes: &[u8],
id: &RequestId,
maximum: usize,
) -> Result<(JsonRpcResponse, String), ManagedTasksError> {
let (response, source) = decode_strict_jsonrpc_response(bytes, maximum)
.map_err(|_| ManagedTasksError::InvalidResponse)?
.into_parts();
if !response
.id
.as_ref()
.is_some_and(|actual| actual.correlates_with(id))
{
return Err(ManagedTasksError::ResponseIdMismatch);
}
if let Some(error) = &response.error {
return Err(ManagedTasksError::Remote {
code: error.code.clone(),
});
}
Ok((response, source.ok_or(ManagedTasksError::InvalidResponse)?))
}
fn admit_composition(
discovery: &CoreRequest,
id: &RequestId,
bytes: &[u8],
decoder: &Decoder,
maximum: usize,
) -> Result<(), ClientCredentialsTasksError> {
admit_resource(discovery, id, bytes)?;
let (response, source) = response_source(bytes, id, maximum)?;
let CoreResult::Final(FinalCoreResult::Discover(result)) = discovery
.decode_response_result(&response, &source)
.map_err(|_| ManagedTasksError::InvalidResponse)?
else {
return Err(ManagedTasksError::Negotiation.into());
};
admit_tasks(&result, decoder)?;
Ok(())
}
fn admit_tasks(
discovery: &ServerDiscoverResult,
decoder: &Decoder,
) -> Result<(), ManagedTasksError> {
match decoder {
Decoder::Tool(_) => crate::admit_final_tasks_result_discriminator(discovery, "task"),
Decoder::Get(_) => crate::admit_final_tasks_discovery_surface(
discovery,
TASK_GET,
ExtensionDirection::ClientToServer,
),
Decoder::Update => crate::admit_final_tasks_discovery_surface(
discovery,
TASK_UPDATE,
ExtensionDirection::ClientToServer,
),
Decoder::Cancel => crate::admit_final_tasks_discovery_surface(
discovery,
TASK_CANCEL,
ExtensionDirection::ClientToServer,
),
}
.map_err(|_| ManagedTasksError::Negotiation)
}
fn decode_result(
decoder: &Decoder,
bytes: &[u8],
id: &RequestId,
maximum: usize,
) -> Result<ManagedTaskEvent, ManagedTasksError> {
let (response, source) = response_source(bytes, id, maximum)?;
match decoder {
Decoder::Tool(request) => {
if response
.result
.as_ref()
.and_then(|result| result.get("resultType"))
.and_then(Value::as_str)
== Some("task")
{
validate_task_input_shapes(&response)?;
}
let CoreResult::Final(result) = request
.decode_response_result(&response, &source)
.map_err(|_| ManagedTasksError::InvalidResponse)?
else {
return Err(ManagedTasksError::InvalidResponse);
};
match result {
FinalCoreResult::ToolsCall { .. }
| FinalCoreResult::ToolsCallTask { .. }
| FinalCoreResult::ToolsCallInputRequired { .. } => {
Ok(ManagedTaskEvent::ToolResult(Box::new(result)))
}
_ => Err(ManagedTasksError::InvalidResponse),
}
}
Decoder::Get(expected) => {
validate_task_input_shapes(&response)?;
let result: GetTaskResult =
serde_json::from_str(&source).map_err(|_| ManagedTasksError::InvalidResponse)?;
if &result.task.base().task_id != expected {
return Err(ManagedTasksError::TaskIdMismatch);
}
Ok(ManagedTaskEvent::Snapshot(Box::new(result)))
}
Decoder::Update => serde_json::from_str(&source)
.map(ManagedTaskEvent::Updated)
.map_err(|_| ManagedTasksError::InvalidResponse),
Decoder::Cancel => serde_json::from_str(&source)
.map(ManagedTaskEvent::Cancelled)
.map_err(|_| ManagedTasksError::InvalidResponse),
}
}
fn decode_record(
decoder: &Decoder,
bytes: &[u8],
id: &RequestId,
maximum: usize,
progress: Option<&ProgressMarker>,
last_progress: &mut Option<ExactNonNegativeJsonNumber>,
) -> Result<ManagedTaskEvent, ManagedTasksError> {
match decode_strict_jsonrpc_message(bytes, maximum)
.map_err(|_| ManagedTasksError::InvalidResponse)?
{
JsonRpcMessage::Response(_) => decode_result(decoder, bytes, id, maximum),
JsonRpcMessage::Request(request) => {
if request.id.is_some() {
return Err(ManagedTasksError::InvalidResponse);
}
#[derive(Deserialize)]
struct Raw {
params: Option<Box<serde_json::value::RawValue>>,
}
let raw: Raw =
serde_json::from_slice(bytes).map_err(|_| ManagedTasksError::InvalidResponse)?;
let notification = match raw.params {
Some(params) => ServerNotification::decode_with_raw_params(&request, params.get()),
None => ServerNotification::decode(&request),
}
.map_err(|_| ManagedTasksError::InvalidResponse)?;
match ¬ification {
ServerNotification::Cancelled(_)
| ServerNotification::SubscriptionsAcknowledged(_) => {
return Err(ManagedTasksError::InvalidResponse);
}
ServerNotification::Progress(update) => {
if progress != Some(&update.progress_token)
|| last_progress
.as_ref()
.is_some_and(|last| update.progress.cmp(last).is_le())
{
return Err(ManagedTasksError::InvalidProgress);
}
*last_progress = Some(update.progress.clone());
}
_ => {}
}
Ok(ManagedTaskEvent::Notification(Box::new(notification)))
}
}
}
#[cfg(test)]
mod tests {
use super::super::rpc::tests::{NativeEnd, native_response_for, runtime};
use super::*;
use fastmcp_protocol::ClientCapabilities;
fn metadata() -> Value {
task_metadata(FinalRequestMeta::new(ClientCapabilities::default())).unwrap()
}
fn id() -> TaskId {
TaskId::parse("owned-task").unwrap()
}
fn task() -> Task {
serde_json::from_value(json!({"taskId":"owned-task", "status":"input_required",
"createdAt":"2026-09-17T00:00:00Z", "lastUpdatedAt":"2026-09-17T00:00:00Z", "ttlMs":60000,
"inputRequests":{"roots":{"method":"roots/list"}}})).unwrap()
}
fn envelope(result: Value) -> Vec<u8> {
serde_json::to_vec(&json!({"jsonrpc":"2.0", "id":2, "result":result})).unwrap()
}
fn prepared(request: ManagedTaskRequest) -> Prepared {
prepare(
"https://resource.example/mcp",
&metadata(),
&RequestId::Number(2),
request,
ClientCredentialsTasksLimits::default(),
)
.unwrap()
}
#[test]
fn both_extensions_are_explicit_and_other_profiles_cannot_be_smuggled_in() {
let meta = metadata();
assert_eq!(
meta[FINAL_CLIENT_CAPABILITIES_META_KEY]["extensions"],
json!({TASKS_EXTENSION:{},CLIENT_CREDENTIALS_EXTENSION:{}})
);
for extensions in [
json!({TASKS_EXTENSION:{"extra":true}}),
json!({CLIENT_CREDENTIALS_EXTENSION:{"extra":true}}),
json!({"com.example/other":{}}),
] {
let mut value =
serde_json::to_value(FinalRequestMeta::new(ClientCapabilities::default())).unwrap();
value[FINAL_CLIENT_CAPABILITIES_META_KEY]["extensions"] = extensions;
assert!(task_metadata(serde_json::from_value(value).unwrap()).is_err());
}
}
#[test]
fn commands_preserve_the_composed_profile_and_use_the_official_wire_verbs() {
for (request, method) in [
(
ManagedTaskRequest::CallTool {
name: "compute".to_owned(),
arguments: Some(json!({"input":1})),
},
"tools/call",
),
(ManagedTaskRequest::Get(id()), TASK_GET),
(ManagedTaskRequest::Cancel(id()), TASK_CANCEL),
(
ManagedTaskRequest::Update {
task: Box::new(task()),
input_responses: serde_json::from_value(json!({"roots":{"roots":[]}})).unwrap(),
},
TASK_UPDATE,
),
] {
let request = prepared(request);
let body: Value = serde_json::from_slice(request.wire.body()).unwrap();
assert_eq!(body["method"], method);
assert_eq!(body["id"], 2);
assert_eq!(
body["params"]["_meta"][FINAL_CLIENT_CAPABILITIES_META_KEY]["extensions"],
json!({TASKS_EXTENSION:{},CLIENT_CREDENTIALS_EXTENSION:{}})
);
assert!(
!request
.wire
.headers()
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization"))
);
}
}
#[test]
fn discovery_requires_tasks_and_validates_optional_auth_and_response_owner() {
let discovery = CoreRequest::decode(
fastmcp_protocol::protocol_policy::ProtocolEra::Modern2026,
"server/discover",
Some(&json!({"_meta":metadata()})),
)
.unwrap();
let base = json!({"resultType":"complete", "supportedVersions":[FINAL_PROTOCOL_VERSION], "ttlMs":0, "cacheScope":"private",
"capabilities":{"extensions":{TASKS_EXTENSION:{},CLIENT_CREDENTIALS_EXTENSION:{}}}});
assert!(
admit_composition(
&discovery,
&RequestId::Number(2),
&envelope(base.clone()),
&Decoder::Get(id()),
65536
)
.is_ok()
);
for key in [TASKS_EXTENSION, CLIENT_CREDENTIALS_EXTENSION] {
let mut changed = base.clone();
changed["capabilities"]["extensions"]
.as_object_mut()
.unwrap()
.remove(key);
let admitted = admit_composition(
&discovery,
&RequestId::Number(2),
&envelope(changed),
&Decoder::Get(id()),
65536,
);
assert_eq!(admitted.is_ok(), key == CLIENT_CREDENTIALS_EXTENSION);
let mut malformed = base.clone();
malformed["capabilities"]["extensions"][key] = Value::Null;
assert!(
admit_composition(
&discovery,
&RequestId::Number(2),
&envelope(malformed),
&Decoder::Get(id()),
65536
)
.is_err()
);
}
assert!(
admit_composition(
&discovery,
&RequestId::Number(3),
&envelope(base),
&Decoder::Get(id()),
65536
)
.is_err()
);
}
#[test]
fn task_and_response_identities_are_independent_and_updates_validate_the_ledger() {
let mut result = serde_json::to_value(task()).unwrap();
result["resultType"] = json!("complete");
assert!(matches!(
decode_result(
&Decoder::Get(id()),
&envelope(result.clone()),
&RequestId::Number(2),
65536
),
Ok(ManagedTaskEvent::Snapshot(_))
));
assert!(matches!(
decode_result(
&Decoder::Get(TaskId::parse("other").unwrap()),
&envelope(result.clone()),
&RequestId::Number(2),
65536
),
Err(ManagedTasksError::TaskIdMismatch)
));
assert!(matches!(
decode_result(
&Decoder::Get(id()),
&envelope(result),
&RequestId::String("2".to_owned()),
65536
),
Err(ManagedTasksError::ResponseIdMismatch)
));
let invalid = ManagedTaskRequest::Update {
task: Box::new(task()),
input_responses: serde_json::from_value(json!({"roots":{"action":"accept"}})).unwrap(),
};
assert!(matches!(
prepare(
"https://resource.example/mcp",
&metadata(),
&RequestId::Number(2),
invalid,
ClientCredentialsTasksLimits::default()
),
Err(ManagedTasksError::InvalidInputResponses)
));
}
#[test]
fn streaming_uses_exact_progress_without_mutating_state_on_rejection() {
let request = prepared(ManagedTaskRequest::CallTool {
name: "compute".to_owned(),
arguments: None,
});
let token = ProgressMarker::String("owned".to_owned());
let mut last = None;
let frame = |token, number| {
serde_json::to_vec(&json!({"jsonrpc":"2.0", "method":"notifications/progress", "params":{"progressToken":token,"progress":number}})).unwrap()
};
assert!(
decode_record(
&request.decoder,
&frame("owned", 1),
&RequestId::Number(2),
4096,
Some(&token),
&mut last
)
.is_ok()
);
let before = last.clone();
for bytes in [frame("other", 2), frame("owned", 1)] {
assert!(matches!(
decode_record(
&request.decoder,
&bytes,
&RequestId::Number(2),
4096,
Some(&token),
&mut last
),
Err(ManagedTasksError::InvalidProgress)
));
assert_eq!(last, before);
}
}
#[test]
fn tool_results_may_be_immediate_input_required_or_a_task_without_coercion() {
let request = prepared(ManagedTaskRequest::CallTool {
name: "compute".to_owned(),
arguments: None,
});
let mut task_result = serde_json::to_value(task()).unwrap();
task_result["resultType"] = json!("task");
for result in [
json!({"resultType":"complete", "content":[]}),
json!({"resultType":"input_required", "requestState":"opaque-state"}),
task_result,
] {
assert!(matches!(
decode_result(
&request.decoder,
&envelope(result),
&RequestId::Number(2),
65536
),
Ok(ManagedTaskEvent::ToolResult(_))
));
}
assert!(
decode_result(
&Decoder::Update,
&envelope(json!({"resultType":"complete"})),
&RequestId::Number(2),
65536
)
.is_ok()
);
assert!(
decode_result(
&Decoder::Cancel,
&envelope(json!({"resultType":"complete"})),
&RequestId::Number(2),
65536
)
.is_ok()
);
}
#[test]
fn bounds_and_peer_errors_do_not_hide_unbounded_work_or_leak_diagnostics() {
assert!(ClientCredentialsTasksLimits::new(1, 1, 1, Duration::from_secs(1)).is_ok());
assert!(ClientCredentialsTasksLimits::new(0, 1, 1, Duration::from_secs(1)).is_err());
assert!(ClientCredentialsTasksLimits::new(1, 1, 1025, Duration::from_secs(1)).is_err());
let tiny = ClientCredentialsTasksLimits::new(1, 4096, 1, Duration::from_secs(1)).unwrap();
assert!(matches!(
prepare(
"https://resource.example/mcp",
&metadata(),
&RequestId::Number(2),
ManagedTaskRequest::Get(id()),
tiny
),
Err(ManagedTasksError::RequestTooLarge)
));
let error = response_source(br#"{"jsonrpc":"2.0","id":2,"error":{"code":-32603,"message":"secret-canary","data":"secret-canary"}}"#, &RequestId::Number(2), 4096).err().unwrap();
assert!(matches!(error, ManagedTasksError::Remote { .. }));
assert!(!format!("{error:?} {error}").contains("secret-canary"));
for bytes in [
br#"[{"jsonrpc":"2.0","id":2,"result":{}}]"#.as_slice(),
br#"{"jsonrpc":"2.0","id":2,"id":2,"result":{}}"#.as_slice(),
] {
assert!(response_source(bytes, &RequestId::Number(2), 4096).is_err());
}
}
#[test]
fn only_get_body_read_failure_retains_observation_interruption() {
assert!(matches!(
json_body_error(
&Decoder::Get(id()),
ModernHttpExecutorError::ResponseBodyReadFailed
),
ManagedTasksError::MissingTerminal
));
let tool = prepared(ManagedTaskRequest::CallTool {
name: "effect".to_owned(),
arguments: None,
});
for decoder in [&Decoder::Update, &Decoder::Cancel, &tool.decoder] {
assert!(matches!(
json_body_error(decoder, ModernHttpExecutorError::ResponseBodyReadFailed),
ManagedTasksError::InvalidResponse
));
}
for error in [
ModernHttpExecutorError::Cancelled,
ModernHttpExecutorError::InvalidTimeoutPolicy,
ModernHttpExecutorError::ResponseBodyTooLarge { maximum_bytes: 1 },
ModernHttpExecutorError::CredentialInPeerError,
ModernHttpExecutorError::UnsupportedContentEncoding,
ModernHttpExecutorError::Redirect { status: 307 },
] {
assert!(matches!(
json_body_error(&Decoder::Get(id()), error),
ManagedTasksError::InvalidResponse
));
}
}
#[test]
fn complete_bad_json_cannot_become_a_recoverable_get_interruption() {
let mut result = serde_json::to_value(task()).unwrap();
result["resultType"] = json!("complete");
let bytes = envelope(result);
assert!(matches!(
decode_result(&Decoder::Get(id()), &bytes, &RequestId::Number(2), 65536),
Ok(ManagedTaskEvent::Snapshot(_))
));
for invalid in [
b"{\"jsonrpc\":".as_slice(),
b"not-json".as_slice(),
br#"{"jsonrpc":"2.0","id":2,"result":{"resultType":"complete"}}"#.as_slice(),
] {
assert!(matches!(
decode_result(&Decoder::Get(id()), invalid, &RequestId::Number(2), 65536),
Err(ManagedTasksError::InvalidResponse)
));
}
assert!(matches!(
decode_result(&Decoder::Get(id()), &bytes, &RequestId::Number(3), 65536),
Err(ManagedTasksError::ResponseIdMismatch)
));
}
fn raw_input_task(result_type: &str, input_requests: &str) -> Vec<u8> {
let result = format!(
r#"{{"resultType":"{result_type}","taskId":"owned-task","status":"input_required","createdAt":"2026-09-17T00:00:00Z","lastUpdatedAt":"2026-09-17T00:00:00Z","ttlMs":60000,"inputRequests":{input_requests}}}"#
);
format!(r#"{{"jsonrpc":"2.0","id":2,"result":{result}}}"#).into_bytes()
}
fn decoded_task(event: ManagedTaskEvent) -> Task {
match event {
ManagedTaskEvent::Snapshot(result) => result.task,
ManagedTaskEvent::ToolResult(result) => match *result {
FinalCoreResult::ToolsCallTask { result } => result.task,
_ => panic!("expected a task-producing tool result"),
},
_ => panic!("expected a task event"),
}
}
#[test]
fn task_input_descriptors_are_validated_before_get_and_tool_decode() {
let get = Decoder::Get(id());
let tool = prepared(ManagedTaskRequest::CallTool {
name: "compute".to_owned(),
arguments: None,
});
let valid = [
r#"{"method":"roots/list"}"#,
r#"{"method":"roots/list","params":{}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"tools":[],"toolChoice":{"mode":"auto"},"includeContext":"allServers"}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"toolChoice":{}}}"#,
r#"{"method":"elicitation/create","params":{"mode":"form","message":"Approve?","requestedSchema":{"type":"object","properties":{}}}}"#,
r#"{"method":"elicitation/create","params":{"mode":"url","message":"Approve?","url":"https://approval.example/confirm"}}"#,
];
let invalid = [
r#"{"method":"roots/list","params":[]}"#,
r#"{"method":"roots/list","params":null}"#,
r#"{"method":"sampling/createMessage","params":[[],8]}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"tools":null}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"tools":{}}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"toolChoice":null}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"toolChoice":[]}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"toolChoice":{"mode":null}}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"toolChoice":{"mode":{"auto":null}}}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"includeContext":null}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"includeContext":[]}}"#,
r#"{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"includeContext":{"allServers":null}}}"#,
r#"{"method":"elicitation/create","params":["form","Approve?",{"type":"object","properties":{}}]}"#,
];
for (decoder, result_type) in [(&get, "complete"), (&tool.decoder, "task")] {
for descriptor in valid {
let requests = format!(r#"{{"input":{descriptor}}}"#);
let event = decode_result(
decoder,
&raw_input_task(result_type, &requests),
&RequestId::Number(2),
65536,
)
.unwrap_or_else(|error| panic!("valid {descriptor}: {error}"));
let Task::InputRequired {
base,
input_requests,
} = decoded_task(event)
else {
panic!("expected an input-required task");
};
assert_eq!(base.task_id, id());
assert_eq!(input_requests.len(), 1);
assert!(input_requests.contains_key("input"));
}
for descriptor in invalid {
let requests = format!(r#"{{"input":{descriptor}}}"#);
assert!(
matches!(
decode_result(
decoder,
&raw_input_task(result_type, &requests),
&RequestId::Number(2),
65536
),
Err(ManagedTasksError::InvalidResponse)
),
"invalid descriptor was admitted: {descriptor}"
);
}
}
}
#[test]
fn task_input_admission_preserves_large_integers_and_correlation_precedence() {
let get = Decoder::Get(id());
let tool = prepared(ManagedTaskRequest::CallTool {
name: "compute".to_owned(),
arguments: None,
});
for (decoder, result_type) in [(&get, "complete"), (&tool.decoder, "task")] {
let bytes = raw_input_task(
result_type,
r#"{"sample":{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":18446744073709551617}}}"#,
);
let event = decode_result(decoder, &bytes, &RequestId::Number(2), 65536).unwrap();
let Task::InputRequired { input_requests, .. } = decoded_task(event) else {
panic!("expected an input-required task");
};
let fastmcp_protocol::FinalEmbeddedInputRequest::Sampling(params) =
&input_requests["sample"]
else {
panic!("expected a sampling descriptor");
};
assert_eq!(params.max_tokens.as_str(), "18446744073709551617");
let malformed = raw_input_task(
result_type,
r#"{"sample":{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":18446744073709551617,"toolChoice":null}}}"#,
);
assert!(matches!(
decode_result(decoder, &malformed, &RequestId::Number(2), 65536),
Err(ManagedTasksError::InvalidResponse)
));
for wrong_id in [RequestId::Number(3), RequestId::String("2".to_owned())] {
assert!(matches!(
decode_result(decoder, &malformed, &wrong_id, 65536),
Err(ManagedTasksError::ResponseIdMismatch)
));
}
}
}
#[test]
fn task_input_admission_keeps_duplicate_rejection_and_all_sibling_checks() {
let get = Decoder::Get(id());
let tool = prepared(ManagedTaskRequest::CallTool {
name: "compute".to_owned(),
arguments: None,
});
for (decoder, result_type) in [(&get, "complete"), (&tool.decoder, "task")] {
let valid = raw_input_task(
result_type,
r#"{"a-roots":{"method":"roots/list"},"z-sample":{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"metadata":{"nested":{"key":1}}}}}"#,
);
let Task::InputRequired { input_requests, .. } =
decoded_task(decode_result(decoder, &valid, &RequestId::Number(2), 65536).unwrap())
else {
panic!("expected an input-required task");
};
assert_eq!(input_requests.len(), 2);
for requests in [
r#"{"a-roots":{"method":"roots/list"},"z-sample":{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"metadata":{"nested":{"key":1,"key":2}}}}}"#,
r#"{"a-roots":{"method":"roots/list"},"z-sample":{"method":"sampling/createMessage","params":{"messages":[],"maxTokens":8,"metadata":{"nested":{"key":1}},"tools":null}}}"#,
r"[]",
r"null",
] {
assert!(matches!(
decode_result(
decoder,
&raw_input_task(result_type, requests),
&RequestId::Number(2),
65536
),
Err(ManagedTasksError::InvalidResponse)
));
}
assert!(decode_result(decoder, &valid, &RequestId::Number(2), 65536).is_ok());
}
}
#[test]
fn task_input_shape_checks_do_not_reinterpret_completed_results_or_acknowledgements() {
let get = Decoder::Get(id());
let tool = prepared(ManagedTaskRequest::CallTool {
name: "compute".to_owned(),
arguments: None,
});
let nested = json!({
"content": [],
"status": "input_required",
"inputRequests": {"opaque": {"method": "roots/list", "params": []}}
});
for (decoder, result_type) in [(&get, "complete"), (&tool.decoder, "task")] {
let mut result = serde_json::to_value(task()).unwrap();
result.as_object_mut().unwrap().remove("inputRequests");
result["resultType"] = json!(result_type);
result["status"] = json!("completed");
result["result"] = nested.clone();
let completed = decoded_task(
decode_result(decoder, &envelope(result), &RequestId::Number(2), 65536).unwrap(),
);
assert!(matches!(&completed, Task::Completed { .. }));
assert_eq!(serde_json::to_value(completed).unwrap()["result"], nested);
}
let immediate = envelope(json!({
"resultType": "complete",
"content": [],
"status": "input_required",
"inputRequests": []
}));
assert!(matches!(
decode_result(&tool.decoder, &immediate, &RequestId::Number(2), 65536),
Ok(ManagedTaskEvent::ToolResult(result))
if matches!(result.as_ref(), FinalCoreResult::ToolsCall { .. })
));
let acknowledgement = envelope(json!({
"resultType": "complete",
"status": "input_required",
"inputRequests": []
}));
assert!(matches!(
decode_result(
&Decoder::Update,
&acknowledgement,
&RequestId::Number(2),
65536
),
Ok(ManagedTaskEvent::Updated(_))
));
assert!(matches!(
decode_result(
&Decoder::Cancel,
&acknowledgement,
&RequestId::Number(2),
65536
),
Ok(ManagedTaskEvent::Cancelled(_))
));
}
fn finite_tool_results() -> [Value; 3] {
let mut background = serde_json::to_value(task()).unwrap();
background["resultType"] = json!("task");
[
json!({"resultType":"complete","content":[]}),
json!({"resultType":"input_required","requestState":"opaque"}),
background,
]
}
async fn finite_tool_call(
cx: &Cx,
body: &str,
end: NativeEnd,
records: usize,
) -> (ClientCredentialsTaskCall, std::thread::JoinHandle<()>) {
let mut meta = metadata();
meta["progressToken"] = json!("owned-progress");
let prepared = prepare(
"https://machine.example/mcp",
&meta,
&RequestId::Number(7),
ManagedTaskRequest::CallTool {
name: "compute".to_owned(),
arguments: None,
},
ClientCredentialsTasksLimits::default(),
)
.unwrap();
let Decoder::Tool(request) = &prepared.decoder else {
panic!("tool fixture expected");
};
let response =
native_response_for(cx, body, "text/event-stream", end, (**request).clone()).await;
let (
super::super::ClientCredentialsResponse {
response,
snapshot,
owner,
cancellation,
request_id,
deadline,
..
},
peer,
) = response;
let limits =
ClientCredentialsTasksLimits::new(4096, 4096, records, Duration::from_secs(5)).unwrap();
(
ClientCredentialsTaskCall::new(
response,
prepared.decoder,
prepared.progress,
snapshot,
owner,
cancellation,
request_id,
deadline,
limits,
)
.unwrap(),
peer,
)
}
fn finite_tool_frame(result: &Value) -> String {
format!(
"data: {}\n\n",
json!({"jsonrpc":"2.0","id":7,"result":result})
)
}
#[test]
fn native_machine_tasks_preserve_all_result_families_at_clean_eof() {
runtime().block_on(async {
let cx = Cx::current().unwrap();
for (index, result) in finite_tool_results().iter().enumerate() {
let body = format!("{}: final keepalive\n\n", finite_tool_frame(result));
let (mut call, peer) = finite_tool_call(&cx, &body, NativeEnd::Complete, 1).await;
let Some(ManagedTaskEvent::ToolResult(result)) =
call.next_event(&cx).await.unwrap()
else {
panic!("finite tool result expected");
};
assert!(matches!(
(index, &*result),
(0, FinalCoreResult::ToolsCall { .. })
| (1, FinalCoreResult::ToolsCallInputRequired { .. })
| (2, FinalCoreResult::ToolsCallTask { .. })
));
assert_eq!(call.records, 1, "EOF does not spend an extra record");
assert!(call.finished && call.body.is_none());
assert!(call.next_event(&cx).await.unwrap().is_none());
peer.join().unwrap();
}
});
}
#[test]
fn native_machine_tasks_reject_duplicate_or_trailing_records_before_publication() {
runtime().block_on(async {
let cx = Cx::current().unwrap();
for result in finite_tool_results() {
let terminal = finite_tool_frame(&result);
for suffix in [terminal.as_str(), "data: private-malformed-tail\n\n",
"data: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/tools/list_changed\"}\n\n"]
{
let (mut call, peer) = finite_tool_call(&cx, &format!("{terminal}{suffix}"), NativeEnd::Complete, 1).await;
let error = call.next_event(&cx).await.err().expect("trailing record must invalidate terminal");
assert!(matches!(&error, ClientCredentialsTasksError::Protocol(ManagedTasksError::InvalidResponse)));
assert!(!format!("{error:?} {error}").contains("private-malformed-tail"));
assert_eq!(call.records, 0);
assert!(!call.finished && call.body.is_none());
assert!(matches!(call.next_event(&cx).await,
Err(ClientCredentialsTasksError::Protocol(ManagedTasksError::Closed))));
peer.join().unwrap();
}
}
});
}
#[test]
fn native_machine_tasks_reject_incomplete_sse_and_http_tails() {
runtime().block_on(async {
let cx = Cx::current().unwrap();
let terminal = finite_tool_frame(&finite_tool_results()[2]);
for (suffix, end) in [
("data: {}\n", NativeEnd::Complete),
("data: {", NativeEnd::Complete),
("", NativeEnd::Truncate),
] {
let (mut call, peer) =
finite_tool_call(&cx, &format!("{terminal}{suffix}"), end, 1).await;
assert!(call.next_event(&cx).await.is_err());
assert_eq!(call.records, 0);
assert!(!call.finished && call.body.is_none());
peer.join().unwrap();
}
});
}
#[test]
fn native_machine_tasks_keep_delivered_progress_without_publishing_invalid_terminal() {
runtime().block_on(async {
let cx = Cx::current().unwrap();
let progress = "data: {\"jsonrpc\":\"2.0\",\"method\":\"notifications/progress\",\"params\":{\"progressToken\":\"owned-progress\",\"progress\":1}}\n\n";
let terminal = finite_tool_frame(&finite_tool_results()[2]);
let (mut call, peer) = finite_tool_call(&cx,
&format!("{progress}{terminal}{terminal}"), NativeEnd::Complete, 2).await;
assert!(matches!(call.next_event(&cx).await.unwrap(), Some(ManagedTaskEvent::Notification(_))));
let before = call.last_progress.clone();
assert!(before.is_some());
assert!(call.next_event(&cx).await.is_err());
assert_eq!(call.last_progress, before);
assert_eq!(call.records, 1);
assert!(!call.finished);
peer.join().unwrap();
});
}
#[test]
fn native_machine_tasks_abandoned_eof_wait_cannot_publish_a_task() {
runtime().block_on(async {
let cx = Cx::current().unwrap();
let terminal = finite_tool_frame(&finite_tool_results()[2]);
let (mut call, peer) = finite_tool_call(&cx, &terminal, NativeEnd::Hold, 1).await;
let stop = cx.now().saturating_add_nanos(250_000_000);
assert!(
asupersync::time::timeout_at(stop, Box::pin(call.next_event(&cx)))
.await
.is_err()
);
assert_eq!(call.records, 0);
assert!(!call.finished && call.body.is_none());
assert!(matches!(
call.next_event(&cx).await,
Err(ClientCredentialsTasksError::Protocol(
ManagedTasksError::Closed
))
));
peer.join().unwrap();
});
}
#[test]
fn native_machine_tasks_eof_wait_keeps_owner_cancellation_effective() {
use std::future::{Future, poll_fn};
runtime().block_on(async {
let cx = Cx::current().unwrap();
let terminal = finite_tool_frame(&finite_tool_results()[2]);
let (mut call, peer) = finite_tool_call(&cx, &terminal, NativeEnd::Hold, 1).await;
let owner = call.owner.clone();
let mut trigger = Box::pin(asupersync::time::sleep(
cx.now(),
Duration::from_millis(100),
));
let mut pending = Box::pin(call.next_event(&cx));
let mut cancelled = false;
let result = poll_fn(|task| {
if !cancelled && trigger.as_mut().poll(task).is_ready() {
cancelled = true;
owner.cancel();
}
pending.as_mut().poll(task)
})
.await;
drop(pending);
assert!(cancelled && result.is_err());
assert_eq!(call.records, 0);
assert!(!call.finished && call.body.is_none());
peer.join().unwrap();
});
}
}