use crate::Context;
use crate::error::{Error, ErrorCode};
use crate::shared::{ArcSlice, ArcStr};
use crate::types::helpers::TypeCategory;
use crate::types::request::RequestParamsMeta;
use crate::types::{Meta, ProgressToken};
use serde::de::DeserializeOwned;
use serde_json::Value;
use std::collections::HashMap;
#[cfg(feature = "tasks")]
use crate::types::RelatedTaskMetadata;
const POSITIONAL: [&str; 8] = [
"arg0", "arg1", "arg2", "arg3", "arg4", "arg5", "arg6", "arg7",
];
#[inline]
pub(crate) fn positional_name(index: usize) -> &'static str {
match POSITIONAL.get(index) {
Some(name) => name,
None => "",
}
}
#[inline]
pub(crate) fn positional_slot(name: &str) -> Option<usize> {
POSITIONAL.iter().position(|positional| *positional == name)
}
#[derive(Debug, Default, Clone)]
pub struct ArgNames {
names: Option<ArcSlice<ArcStr>>,
arity: usize,
}
impl ArgNames {
#[inline]
pub fn new<I, S>(names: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
let names = names
.into_iter()
.map(|name| ArcStr::from(name.into()))
.collect::<Vec<_>>();
Self {
arity: names.len(),
names: Some(ArcSlice::from(names)),
}
}
#[inline]
pub fn positional(arity: usize) -> Self {
Self { names: None, arity }
}
#[inline]
pub fn is_declared(&self) -> bool {
self.names.is_some()
}
#[inline]
pub fn arity(&self) -> usize {
self.arity
}
#[inline]
pub fn get(&self, index: usize) -> &str {
self.names
.as_deref()
.and_then(|names| names.get(index))
.map_or_else(|| positional_name(index), |name| &**name)
}
#[inline]
pub fn len(&self) -> usize {
self.names.as_deref().map_or(0, <[ArcStr]>::len)
}
#[inline]
pub(crate) fn declare<I, S>(&self, names: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
Self {
arity: self.arity,
..Self::new(names)
}
}
#[inline]
pub(crate) fn slot_of(&self, name: &str) -> Option<usize> {
match self.names.as_deref() {
Some(names) => names.iter().position(|declared| &**declared == name),
None => positional_slot(name),
}
.filter(|slot| *slot < self.arity)
}
#[inline]
pub(crate) fn duplicate(&self) -> Option<&str> {
let names = self.names.as_deref()?;
let mut seen = std::collections::HashSet::with_capacity(names.len());
names
.iter()
.map(|name| &**name)
.find(|name| !seen.insert(*name))
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
pub(crate) enum Payload<'a> {
Args(serde_json::Value),
Meta(&'a Option<RequestParamsMeta>),
}
pub(crate) enum Source {
Args,
Meta,
}
pub(crate) trait RequestArgument: Sized {
type Error;
fn extract(payload: Payload<'_>) -> Result<Self, Self::Error>;
#[inline]
fn source() -> Source {
Source::Args
}
}
pub(crate) trait HandlerArgs {
fn into_parts(self) -> (Option<HashMap<String, Value>>, Option<RequestParamsMeta>);
}
pub trait FromHandlerArgs<P>: Sized {
fn from_args(params: P, names: &ArgNames) -> Result<Self, Error>;
}
impl<'a> Payload<'a> {
#[inline]
pub(crate) fn expect_args(self) -> serde_json::Value {
match self {
Payload::Args(val) => val,
_ => unreachable!("Expected Args variant"),
}
}
#[inline]
pub(crate) fn expect_meta(self) -> &'a Option<RequestParamsMeta> {
match self {
Payload::Meta(meta) => meta,
_ => unreachable!("Expected Meta variant"),
}
}
}
impl<T: DeserializeOwned> RequestArgument for T {
type Error = Error;
#[inline]
fn extract(payload: Payload<'_>) -> Result<Self, Self::Error> {
let arg = payload.expect_args();
T::deserialize(arg).map_err(Error::from)
}
}
impl RequestArgument for Meta<RequestParamsMeta> {
type Error = Error;
#[inline]
fn extract(payload: Payload<'_>) -> Result<Self, Self::Error> {
let meta = payload.expect_meta();
meta.clone()
.ok_or(Error::new(ErrorCode::InvalidParams, "Missing metadata"))
.map(Meta)
}
#[inline]
fn source() -> Source {
Source::Meta
}
}
impl RequestArgument for Meta<ProgressToken> {
type Error = Error;
#[inline]
fn extract(payload: Payload<'_>) -> Result<Self, Self::Error> {
let meta = payload.expect_meta();
meta.as_ref()
.and_then(|meta| meta.progress_token.clone())
.ok_or(Error::new(
ErrorCode::InvalidParams,
"Missing progress token",
))
.map(Meta)
}
#[inline]
fn source() -> Source {
Source::Meta
}
}
#[cfg(feature = "tasks")]
impl RequestArgument for Meta<RelatedTaskMetadata> {
type Error = Error;
#[inline]
fn extract(payload: Payload<'_>) -> Result<Self, Self::Error> {
let meta = payload.expect_meta();
meta.as_ref()
.and_then(|meta| meta.task.clone())
.ok_or(Error::new(
ErrorCode::InvalidParams,
"Missing progress token",
))
.map(Meta)
}
#[inline]
fn source() -> Source {
Source::Meta
}
}
impl RequestArgument for Context {
type Error = Error;
#[inline]
fn extract(payload: Payload<'_>) -> Result<Self, Self::Error> {
let meta = payload.expect_meta();
meta.as_ref()
.and_then(|meta| meta.context.clone())
.ok_or(Error::new(ErrorCode::InvalidParams, "Missing MCP context"))
}
#[inline]
fn source() -> Source {
Source::Meta
}
}
#[inline]
pub(crate) fn extract_arg<T: RequestArgument<Error = Error> + TypeCategory>(
meta: &Option<RequestParamsMeta>,
args: Option<&HashMap<String, Value>>,
names: &ArgNames,
slot: &mut usize,
) -> Result<T, Error> {
match T::source() {
Source::Meta => T::extract(Payload::Meta(meta)),
Source::Args => {
let name = names.get(*slot);
*slot += 1;
match args.and_then(|args| args.get(name)) {
Some(value) => T::extract(Payload::Args(value.clone())).map_err(|err| {
Error::new(
ErrorCode::InvalidParams,
format!("invalid value for argument `{name}`: {err}"),
)
}),
None if T::is_optional() => T::extract(Payload::Args(Value::Null)),
None => Err(Error::new(
ErrorCode::InvalidParams,
format!("missing required argument `{name}`"),
)),
}
}
}
}
impl<P: HandlerArgs> FromHandlerArgs<P> for () {
#[inline]
fn from_args(_: P, _: &ArgNames) -> Result<Self, Error> {
Ok(())
}
}
macro_rules! impl_from_handler_args {
($($T:ident),+) => {
impl<P: HandlerArgs, $($T: RequestArgument<Error = Error> + TypeCategory),+> FromHandlerArgs<P> for ($($T,)+) {
#[inline]
fn from_args(params: P, names: &ArgNames) -> Result<Self, Error> {
let (args, meta) = params.into_parts();
let args = args.as_ref();
let mut slot = 0;
let tuple = (
$(
extract_arg::<$T>(&meta, args, names, &mut slot)?,
)+
);
Ok(tuple)
}
}
};
}
impl_from_handler_args! { T1 }
impl_from_handler_args! { T1, T2 }
impl_from_handler_args! { T1, T2, T3 }
impl_from_handler_args! { T1, T2, T3, T4 }
impl_from_handler_args! { T1, T2, T3, T4, T5 }
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn it_returns_declared_names() {
let names = ArgNames::new(["age", "name"]);
assert_eq!(names.get(0), "age");
assert_eq!(names.get(1), "name");
assert_eq!(names.len(), 2);
assert!(!names.is_empty());
}
#[test]
fn it_falls_back_to_positional_names() {
let names = ArgNames::default();
assert_eq!(names.get(0), "arg0");
assert_eq!(names.get(4), "arg4");
assert!(names.is_empty());
}
#[test]
fn it_falls_back_past_the_declared_names() {
let names = ArgNames::new(["age"]);
assert_eq!(names.get(0), "age");
assert_eq!(names.get(1), "arg1");
}
}