use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use super::{
CallToolResult, CreateMessageRequest, CreateTaskResult, ElicitRequest, GetPromptResult,
ListRootsRequest, MetaObject, ReadResourceResult, ResultType, ServerResult,
};
pub const DEFAULT_MRTR_MAX_ROUNDS: usize = 10;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[non_exhaustive]
pub enum InputRequest {
CreateMessage(CreateMessageRequest),
Elicitation(ElicitRequest),
ListRoots(ListRootsRequest),
}
impl PartialEq for InputRequest {
fn eq(&self, other: &Self) -> bool {
match (serde_json::to_value(self), serde_json::to_value(other)) {
(Ok(a), Ok(b)) => a == b,
_ => false,
}
}
}
pub type InputRequests = BTreeMap<String, InputRequest>;
pub type InputResponses = BTreeMap<String, Value>;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum CallToolResponse {
Complete(CallToolResult),
InputRequired(InputRequiredResult),
Task(CreateTaskResult),
}
impl From<CallToolResult> for CallToolResponse {
fn from(result: CallToolResult) -> Self {
Self::Complete(result)
}
}
impl From<InputRequiredResult> for CallToolResponse {
fn from(result: InputRequiredResult) -> Self {
Self::InputRequired(result)
}
}
impl From<CallToolResponse> for ServerResult {
fn from(response: CallToolResponse) -> Self {
match response {
CallToolResponse::Complete(result) => ServerResult::CallToolResult(result),
CallToolResponse::InputRequired(result) => ServerResult::InputRequiredResult(result),
CallToolResponse::Task(result) => ServerResult::CreateTaskResult(result),
}
}
}
impl From<CreateTaskResult> for CallToolResponse {
fn from(result: CreateTaskResult) -> Self {
Self::Task(result)
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum GetPromptResponse {
Complete(GetPromptResult),
InputRequired(InputRequiredResult),
}
impl From<GetPromptResult> for GetPromptResponse {
fn from(result: GetPromptResult) -> Self {
Self::Complete(result)
}
}
impl From<InputRequiredResult> for GetPromptResponse {
fn from(result: InputRequiredResult) -> Self {
Self::InputRequired(result)
}
}
impl From<GetPromptResponse> for ServerResult {
fn from(response: GetPromptResponse) -> Self {
match response {
GetPromptResponse::Complete(result) => ServerResult::GetPromptResult(result),
GetPromptResponse::InputRequired(result) => ServerResult::InputRequiredResult(result),
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum ReadResourceResponse {
Complete(ReadResourceResult),
InputRequired(InputRequiredResult),
}
impl From<ReadResourceResult> for ReadResourceResponse {
fn from(result: ReadResourceResult) -> Self {
Self::Complete(result)
}
}
impl From<InputRequiredResult> for ReadResourceResponse {
fn from(result: InputRequiredResult) -> Self {
Self::InputRequired(result)
}
}
impl From<ReadResourceResponse> for ServerResult {
fn from(response: ReadResourceResponse) -> Self {
match response {
ReadResourceResponse::Complete(result) => ServerResult::ReadResourceResult(result),
ReadResourceResponse::InputRequired(result) => {
ServerResult::InputRequiredResult(result)
}
}
}
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[non_exhaustive]
pub struct InputRequiredResult {
pub result_type: ResultType,
#[serde(skip_serializing_if = "Option::is_none")]
pub input_requests: Option<InputRequests>,
#[serde(skip_serializing_if = "Option::is_none")]
pub request_state: Option<String>,
#[serde(rename = "_meta", skip_serializing_if = "Option::is_none")]
pub meta: Option<MetaObject>,
}
impl<'de> Deserialize<'de> for InputRequiredResult {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct Helper {
result_type: Option<ResultType>,
input_requests: Option<InputRequests>,
request_state: Option<String>,
#[serde(rename = "_meta")]
meta: Option<MetaObject>,
}
let helper = Helper::deserialize(deserializer)?;
match &helper.result_type {
Some(rt) if rt.is_input_required() => {}
_ => {
return Err(serde::de::Error::custom(
"InputRequiredResult requires resultType to be \"input_required\"",
));
}
}
Ok(InputRequiredResult {
result_type: ResultType::INPUT_REQUIRED,
input_requests: helper.input_requests,
request_state: helper.request_state,
meta: helper.meta,
})
}
}
impl InputRequiredResult {
pub fn new(input_requests: Option<InputRequests>, request_state: Option<String>) -> Self {
Self {
result_type: ResultType::INPUT_REQUIRED,
input_requests,
request_state,
meta: None,
}
}
pub fn from_input_requests(input_requests: InputRequests) -> Self {
Self::new(Some(input_requests), None)
}
pub fn from_request_state(request_state: impl Into<String>) -> Self {
Self::new(None, Some(request_state.into()))
}
pub fn with_meta(mut self, meta: MetaObject) -> Self {
self.meta = Some(meta);
self
}
}
#[cfg(test)]
mod tests {
use super::*;
mod result_type {
use super::*;
#[test]
fn default_is_complete() {
assert_eq!(ResultType::default(), ResultType::COMPLETE);
}
#[test]
fn serializes_complete() {
assert_eq!(
serde_json::to_value(&ResultType::COMPLETE).unwrap(),
serde_json::json!("complete")
);
}
#[test]
fn serializes_input_required() {
assert_eq!(
serde_json::to_value(&ResultType::INPUT_REQUIRED).unwrap(),
serde_json::json!("input_required")
);
}
#[test]
fn deserializes_known_values() {
let complete: ResultType =
serde_json::from_value(serde_json::json!("complete")).unwrap();
assert_eq!(complete, ResultType::COMPLETE);
let input_required: ResultType =
serde_json::from_value(serde_json::json!("input_required")).unwrap();
assert_eq!(input_required, ResultType::INPUT_REQUIRED);
}
#[test]
fn preserves_unknown_extension_values() {
let custom: ResultType =
serde_json::from_value(serde_json::json!("streaming")).unwrap();
assert_eq!(custom.as_str(), "streaming");
assert!(!custom.is_complete());
assert!(!custom.is_input_required());
let reserialized = serde_json::to_value(&custom).unwrap();
assert_eq!(reserialized, serde_json::json!("streaming"));
}
}
mod input_required_result {
use super::*;
#[test]
fn deserializes_with_requests_and_state() {
let json = serde_json::json!({
"resultType": "input_required",
"inputRequests": {
"github_login": {
"method": "elicitation/create",
"params": {
"message": "Please provide your GitHub username",
"requestedSchema": {
"type": "object",
"properties": { "name": { "type": "string" } },
"required": ["name"]
}
}
},
"capital_of_france": {
"method": "sampling/createMessage",
"params": {
"messages": [{
"role": "user",
"content": { "type": "text", "text": "What is the capital of France?" }
}],
"maxTokens": 100
}
}
},
"requestState": "eyJsb2NhdGlvbiI6Ik5ldyBZb3JrIn0"
});
let result: InputRequiredResult = serde_json::from_value(json).unwrap();
let requests = result
.input_requests
.as_ref()
.expect("should have input_requests");
assert_eq!(requests.len(), 2);
assert!(requests.contains_key("github_login"));
assert!(requests.contains_key("capital_of_france"));
assert_eq!(
result.request_state.as_deref(),
Some("eyJsb2NhdGlvbiI6Ik5ldyBZb3JrIn0")
);
}
#[test]
fn roundtrip_preserves_all_fields() {
let json = serde_json::json!({
"resultType": "input_required",
"inputRequests": {
"key": {
"method": "elicitation/create",
"params": {
"message": "test",
"requestedSchema": { "type": "object", "properties": {} }
}
}
},
"requestState": "abc123"
});
let result: InputRequiredResult = serde_json::from_value(json).unwrap();
let reserialized = serde_json::to_value(&result).unwrap();
assert_eq!(reserialized["resultType"], "input_required");
assert!(reserialized["inputRequests"].is_object());
assert_eq!(reserialized["requestState"], "abc123");
}
#[test]
fn deserializes_with_request_state_only() {
let json = serde_json::json!({
"resultType": "input_required",
"requestState": "eyJwcm9ncmVzcyI6IjUwJSJ9"
});
let result: InputRequiredResult = serde_json::from_value(json).unwrap();
assert!(result.input_requests.is_none());
assert_eq!(
result.request_state.as_deref(),
Some("eyJwcm9ncmVzcyI6IjUwJSJ9")
);
}
#[test]
fn rejects_missing_result_type() {
let json = serde_json::json!({
"requestState": "some-state"
});
let err = serde_json::from_value::<InputRequiredResult>(json).unwrap_err();
assert!(
err.to_string().contains("input_required"),
"error should mention the required resultType, got: {err}"
);
}
#[test]
fn rejects_wrong_result_type() {
let json = serde_json::json!({
"resultType": "complete",
"requestState": "some-state"
});
let err = serde_json::from_value::<InputRequiredResult>(json).unwrap_err();
assert!(
err.to_string().contains("input_required"),
"error should mention the required resultType, got: {err}"
);
}
}
mod input_responses {
use super::*;
#[test]
fn deserializes_heterogeneous_results() {
let json = serde_json::json!({
"github_login": {
"action": "accept",
"content": { "name": "octocat" }
},
"capital_of_france": {
"role": "assistant",
"content": { "type": "text", "text": "Paris." },
"model": "claude-3-sonnet-20240307",
"stopReason": "endTurn"
}
});
let responses: InputResponses = serde_json::from_value(json).unwrap();
assert_eq!(responses.len(), 2);
assert!(responses.contains_key("github_login"));
assert!(responses.contains_key("capital_of_france"));
}
}
mod constructors {
use super::*;
#[test]
fn from_request_state_sets_state_only() {
let result = InputRequiredResult::from_request_state("opaque");
assert_eq!(result.result_type, ResultType::INPUT_REQUIRED);
assert!(result.input_requests.is_none());
assert_eq!(result.request_state.as_deref(), Some("opaque"));
}
#[test]
fn from_input_requests_sets_requests_only() {
let mut requests = InputRequests::new();
requests.insert(
"key".to_string(),
serde_json::from_value(serde_json::json!({
"method": "elicitation/create",
"params": {
"message": "test",
"requestedSchema": { "type": "object", "properties": {} }
}
}))
.unwrap(),
);
let result = InputRequiredResult::from_input_requests(requests);
assert!(result.input_requests.is_some());
assert!(result.request_state.is_none());
}
}
}