use super::*;
pub(super) type TaskOwnerResolver =
Arc<dyn Fn(&crate::context::Extensions) -> TaskOwnerResolution + Send + Sync + 'static>;
pub(super) enum TaskOwnerResolution {
Resolved(crate::async_task::TaskOwner),
Invalid,
}
impl TaskOwnerResolution {
pub(super) fn into_owner(self) -> Option<crate::async_task::TaskOwner> {
match self {
Self::Resolved(owner) => Some(owner),
Self::Invalid => None,
}
}
pub(super) fn matches(&self, owner: &crate::async_task::TaskOwner) -> bool {
match self {
Self::Resolved(principal) => {
crate::async_task::owner_matches(owner, principal.as_deref())
}
Self::Invalid => false,
}
}
}
pub(super) fn default_task_owner_resolver() -> TaskOwnerResolver {
Arc::new(|extensions| TaskOwnerResolution::Resolved(oauth_task_owner(extensions)))
}
pub(super) fn custom_task_owner_resolver<F>(resolver: F) -> TaskOwnerResolver
where
F: Fn(&crate::context::Extensions) -> Option<String> + Send + Sync + 'static,
{
Arc::new(move |extensions| {
let resolved =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| resolver(extensions)));
match resolved {
Ok(Some(owner)) if !owner.trim().is_empty() => {
TaskOwnerResolution::Resolved(Some(owner))
}
Ok(Some(_)) => {
tracing::error!(
target: "mcp::tasks",
"task owner resolver returned an empty owner; denying the operation"
);
TaskOwnerResolution::Invalid
}
Ok(None) => TaskOwnerResolution::Resolved(None),
Err(_) => {
tracing::error!(
target: "mcp::tasks",
"task owner resolver panicked; denying the operation"
);
TaskOwnerResolution::Invalid
}
}
})
}
#[cfg(feature = "oauth")]
fn oauth_task_owner(extensions: &crate::context::Extensions) -> Option<String> {
extensions
.get::<crate::oauth::token::TokenClaims>()
.and_then(|claims| claims.sub.clone())
}
#[cfg(not(feature = "oauth"))]
fn oauth_task_owner(_extensions: &crate::context::Extensions) -> Option<String> {
None
}
#[cfg(feature = "stateless")]
pub(super) fn final_client_capabilities(
extensions: &crate::context::Extensions,
) -> Option<&ClientCapabilities> {
extensions
.get::<crate::stateless::StatelessRequestMeta>()
.and_then(|meta| meta.client_capabilities.as_ref())
}
#[cfg(not(feature = "stateless"))]
pub(super) fn final_client_capabilities(
_extensions: &crate::context::Extensions,
) -> Option<&ClientCapabilities> {
None
}
#[cfg(feature = "stateless")]
pub(super) fn json_value_contains(
actual: &serde_json::Value,
required: &serde_json::Value,
) -> bool {
match (actual, required) {
(serde_json::Value::Object(actual), serde_json::Value::Object(required)) => {
required.iter().all(|(key, value)| {
actual
.get(key)
.is_some_and(|a| json_value_contains(a, value))
})
}
_ => actual == required,
}
}
#[cfg(feature = "stateless")]
pub(super) fn client_capabilities_satisfy(
actual: &ClientCapabilities,
required: &ClientCapabilities,
) -> bool {
let actual = serde_json::to_value(actual).expect("ClientCapabilities is always serializable");
let mut required =
serde_json::to_value(required).expect("ClientCapabilities is always serializable");
if required.pointer("/roots/listChanged") == Some(&serde_json::Value::Bool(false))
&& let Some(roots) = required
.get_mut("roots")
.and_then(serde_json::Value::as_object_mut)
{
roots.remove("listChanged");
}
json_value_contains(&actual, &required)
}