use serde::{de, Deserialize, Deserializer, Serialize};
use serde_json::Value;
use std::collections::HashMap;
#[derive(Debug, Clone, Serialize)]
#[serde(tag = "mode")]
pub enum ElicitRequestParams {
#[serde(rename = "form", rename_all = "camelCase")]
Form {
message: String,
requested_schema: Value,
},
#[serde(rename = "url", rename_all = "camelCase")]
Url {
message: String,
elicitation_id: String,
url: String,
},
}
const ELICIT_MODE_FORM: &str = "form";
const ELICIT_MODE_URL: &str = "url";
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct FormShape {
message: String,
requested_schema: Value,
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct UrlShape {
message: String,
elicitation_id: String,
url: String,
}
impl<'de> Deserialize<'de> for ElicitRequestParams {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let raw = Value::deserialize(deserializer)?;
let mode = match raw.get("mode") {
None | Some(Value::Null) => ELICIT_MODE_FORM,
Some(Value::String(mode)) => mode.as_str(),
Some(_) => return Err(de::Error::custom("`mode` must be a string")),
};
match mode {
ELICIT_MODE_FORM => {
let shape = FormShape::deserialize(&raw).map_err(de::Error::custom)?;
Ok(Self::Form {
message: shape.message,
requested_schema: shape.requested_schema,
})
},
ELICIT_MODE_URL => {
let shape = UrlShape::deserialize(&raw).map_err(de::Error::custom)?;
Ok(Self::Url {
message: shape.message,
elicitation_id: shape.elicitation_id,
url: shape.url,
})
},
other => Err(de::Error::unknown_variant(
other,
&[ELICIT_MODE_FORM, ELICIT_MODE_URL],
)),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ElicitResult {
pub action: ElicitAction,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<HashMap<String, Value>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub enum ElicitAction {
Accept,
Decline,
Cancel,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ElicitationCompleteNotification {
pub elicitation_id: String,
pub result: ElicitResult,
}
#[deprecated(since = "2.0.0", note = "Use ElicitRequestParams instead")]
pub type ElicitInputRequest = ElicitRequestParams;
#[deprecated(since = "2.0.0", note = "Use ElicitResult instead")]
pub type ElicitInputResponse = ElicitResult;
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn elicit_request_form_mode_serialization() {
let params = ElicitRequestParams::Form {
message: "Enter your name".to_string(),
requested_schema: json!({
"type": "object",
"properties": {
"name": { "type": "string" }
}
}),
};
let json = serde_json::to_value(¶ms).unwrap();
assert_eq!(json["mode"], "form");
assert_eq!(json["message"], "Enter your name");
assert!(json["requestedSchema"]["properties"]["name"].is_object());
let roundtrip: ElicitRequestParams = serde_json::from_value(json).unwrap();
match roundtrip {
ElicitRequestParams::Form { message, .. } => {
assert_eq!(message, "Enter your name");
},
ElicitRequestParams::Url { .. } => panic!("Expected Form variant"),
}
}
#[test]
fn elicit_request_url_mode_serialization() {
let params = ElicitRequestParams::Url {
message: "Please authenticate".to_string(),
elicitation_id: "auth-123".to_string(),
url: "https://example.com/auth".to_string(),
};
let json = serde_json::to_value(¶ms).unwrap();
assert_eq!(json["mode"], "url");
assert_eq!(json["message"], "Please authenticate");
assert_eq!(json["elicitationId"], "auth-123");
assert_eq!(json["url"], "https://example.com/auth");
let roundtrip: ElicitRequestParams = serde_json::from_value(json).unwrap();
match roundtrip {
ElicitRequestParams::Url { elicitation_id, .. } => {
assert_eq!(elicitation_id, "auth-123");
},
ElicitRequestParams::Form { .. } => panic!("Expected Url variant"),
}
}
#[test]
fn elicit_request_form_mode_is_optional() {
let params: ElicitRequestParams = serde_json::from_value(json!({
"message": "What is your name?",
"requestedSchema": { "type": "object" }
}))
.expect("a mode-less form elicitation must deserialize");
match params {
ElicitRequestParams::Form {
message,
requested_schema,
} => {
assert_eq!(message, "What is your name?");
assert_eq!(requested_schema["type"], "object");
},
ElicitRequestParams::Url { .. } => panic!("Expected Form variant"),
}
}
#[test]
fn elicit_request_explicit_form_mode_still_deserializes() {
let params: ElicitRequestParams = serde_json::from_value(json!({
"mode": "form",
"message": "hi",
"requestedSchema": {}
}))
.expect("an explicit form mode must still deserialize");
assert!(matches!(params, ElicitRequestParams::Form { .. }));
}
#[test]
fn elicit_request_url_mode_still_requires_its_fields() {
let params: ElicitRequestParams = serde_json::from_value(json!({
"mode": "url",
"message": "auth",
"elicitationId": "auth-1",
"url": "https://example.com"
}))
.expect("a complete url elicitation must deserialize");
assert!(matches!(params, ElicitRequestParams::Url { .. }));
assert!(serde_json::from_value::<ElicitRequestParams>(
json!({ "mode": "url", "message": "auth", "url": "https://example.com" })
)
.is_err());
assert!(serde_json::from_value::<ElicitRequestParams>(
json!({ "mode": "url", "message": "auth", "elicitationId": "auth-1" })
)
.is_err());
}
#[test]
fn elicit_request_rejects_an_unknown_mode() {
assert!(serde_json::from_value::<ElicitRequestParams>(json!({
"mode": "bogus",
"message": "hi",
"requestedSchema": {}
}))
.is_err());
}
#[test]
fn elicit_request_rejects_a_non_string_mode() {
assert!(serde_json::from_value::<ElicitRequestParams>(json!({
"mode": 7,
"message": "hi",
"requestedSchema": {}
}))
.is_err());
}
#[test]
fn elicit_request_form_still_serializes_the_mode_tag() {
let params = ElicitRequestParams::Form {
message: "hi".to_string(),
requested_schema: json!({}),
};
let value = serde_json::to_value(¶ms).unwrap();
assert_eq!(value["mode"], "form");
assert_eq!(value["message"], "hi");
assert!(value["requestedSchema"].is_object());
assert_eq!(
serde_json::to_string(¶ms).unwrap(),
r#"{"mode":"form","message":"hi","requestedSchema":{}}"#
);
}
#[test]
fn elicit_request_form_missing_required_fields_is_an_error() {
assert!(serde_json::from_value::<ElicitRequestParams>(json!({ "message": "hi" })).is_err());
assert!(serde_json::from_value::<ElicitRequestParams>(json!({})).is_err());
}
#[test]
fn elicit_result_accept() {
let mut content = HashMap::new();
content.insert("name".to_string(), json!("Alice"));
let result = ElicitResult {
action: ElicitAction::Accept,
content: Some(content),
};
let json = serde_json::to_value(&result).unwrap();
assert_eq!(json["action"], "accept");
assert_eq!(json["content"]["name"], "Alice");
let roundtrip: ElicitResult = serde_json::from_value(json).unwrap();
assert_eq!(roundtrip.action, ElicitAction::Accept);
assert!(roundtrip.content.is_some());
}
#[test]
fn elicit_result_decline() {
let result = ElicitResult {
action: ElicitAction::Decline,
content: None,
};
let json = serde_json::to_value(&result).unwrap();
assert_eq!(json["action"], "decline");
assert!(json.get("content").is_none());
}
#[test]
fn elicit_action_values() {
assert_eq!(
serde_json::to_value(ElicitAction::Accept).unwrap(),
"accept"
);
assert_eq!(
serde_json::to_value(ElicitAction::Decline).unwrap(),
"decline"
);
assert_eq!(
serde_json::to_value(ElicitAction::Cancel).unwrap(),
"cancel"
);
}
}