Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
124 changes: 118 additions & 6 deletions rust/crates/sift_mcp/src/client_event.rs
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
use std::{collections::HashMap, sync::LazyLock, time::Duration};
use std::{collections::HashMap, path::Path, sync::LazyLock, time::Duration};

use reqwest::header::USER_AGENT;
use rmcp::model::JsonObject;
use serde::Serialize;
use serde_json::Value;

use crate::ClientName;

Expand Down Expand Up @@ -42,6 +44,8 @@ struct ClientEventTarget {
#[derive(Serialize)]
struct ClientEventRequest {
event: &'static str,
#[serde(skip_serializing_if = "HashMap::is_empty")]
properties: HashMap<&'static str, String>,
}

impl ClientEventReporter {
Expand Down Expand Up @@ -79,7 +83,11 @@ impl ClientEventReporter {
self.target.is_some()
}

pub(crate) async fn send(&self, tool_name: &str) -> reqwest::Result<()> {
pub(crate) async fn send(
&self,
tool_name: &str,
arguments: Option<&JsonObject>,
) -> reqwest::Result<()> {
let Some(target) = &self.target else {
return Ok(());
};
Expand All @@ -93,7 +101,10 @@ impl ClientEventReporter {
.timeout(REQUEST_TIMEOUT)
.bearer_auth(&target.api_key)
.header(USER_AGENT, &target.user_agent)
.json(&ClientEventRequest { event })
.json(&ClientEventRequest {
event,
properties: properties_for_tool(tool_name, arguments),
})
.send()
.await?
.error_for_status()?;
Expand All @@ -106,6 +117,42 @@ pub(crate) fn event_for_tool(tool_name: &str) -> Option<&'static str> {
TOOL_EVENTS.get(tool_name).map(String::as_str)
}

/// Properties read from a call's arguments. The event fires before the tool
/// runs, so they describe the request, not its result.
fn properties_for_tool(
tool_name: &str,
arguments: Option<&JsonObject>,
) -> HashMap<&'static str, String> {
let mut properties = HashMap::new();
if tool_name == "create_artifact"
&& let Some(file_type) = arguments.and_then(artifact_file_type)
{
properties.insert("file_type", file_type);
}
properties
}

/// `json` for a structured artifact, otherwise the lowercase extension of
/// `file_path`. An artifact created without a file has no file type.
fn artifact_file_type(arguments: &JsonObject) -> Option<String> {
let storage_class = arguments
.get("storage_class")
.and_then(Value::as_str)
.map(|value| value.trim().to_ascii_lowercase());
if matches!(
storage_class.as_deref(),
Some("structured" | "artifact_storage_class_structured")
) {
return Some("json".to_string());
}
let file_path = arguments.get("file_path").and_then(Value::as_str)?;
Path::new(file_path.trim())
.extension()
.and_then(|extension| extension.to_str())
.filter(|extension| !extension.is_empty())
.map(str::to_ascii_lowercase)
}

#[cfg(test)]
pub(crate) async fn start_http_server(
response: Vec<u8>,
Expand Down Expand Up @@ -167,9 +214,74 @@ pub(crate) async fn start_event_server() -> (String, tokio::task::JoinHandle<Vec

#[cfg(test)]
mod tests {
use super::{ClientEventConfig, ClientEventReporter, event_for_tool, start_event_server};
use super::{
ClientEventConfig, ClientEventReporter, event_for_tool, properties_for_tool,
start_event_server,
};
use crate::ClientName;

fn arguments(value: serde_json::Value) -> rmcp::model::JsonObject {
value.as_object().unwrap().clone()
}

#[test]
fn create_artifact_reports_the_file_type() {
let cases = [
(
serde_json::json!({ "file_path": "/workspace/out/Report.PDF" }),
Some("pdf"),
),
(
serde_json::json!({ "file_path": "dashboard.html" }),
Some("html"),
),
(
serde_json::json!({ "storage_class": "Structured", "payload": {} }),
Some("json"),
),
(serde_json::json!({ "file_path": "notes" }), None),
(serde_json::json!({ "title": "No file yet" }), None),
];
for (args, expected) in cases {
let properties = properties_for_tool("create_artifact", Some(&arguments(args.clone())));
assert_eq!(
properties.get("file_type").map(String::as_str),
expected,
"{args}"
);
}
assert!(properties_for_tool("create_artifact", None).is_empty());
}

#[test]
fn other_tools_report_no_properties() {
let args = arguments(serde_json::json!({ "file_path": "report.pdf" }));
assert!(properties_for_tool("update_artifact", Some(&args)).is_empty());
}

#[tokio::test]
async fn sends_create_artifact_file_type_as_a_property() {
let (rest_uri, server) = start_event_server().await;
let reporter = ClientEventReporter::new(
ClientEventConfig::new(rest_uri, "test-key".to_string()),
ClientName::Chat,
"7.8.9",
);

let args = arguments(serde_json::json!({ "file_path": "report.pdf" }));
reporter.send("create_artifact", Some(&args)).await.unwrap();
let request = String::from_utf8(server.await.unwrap()).unwrap();
let (_, body) = request.split_once("\r\n\r\n").unwrap();

assert_eq!(
serde_json::from_str::<serde_json::Value>(body).unwrap(),
serde_json::json!({
"event": "CLIENT_EVENT_USER_CALLED_MCP_TOOL_CREATE_ARTIFACT",
"properties": { "file_type": "pdf" }
})
);
}

#[test]
fn artifact_archive_tools_have_client_events() {
assert_eq!(
Expand All @@ -195,7 +307,7 @@ mod tests {
"7.8.9",
);

reporter.send("list_assets").await.unwrap();
reporter.send("list_assets", None).await.unwrap();
let request = String::from_utf8(server.await.unwrap()).unwrap();
let (headers, body) = request.split_once("\r\n\r\n").unwrap();

Expand Down Expand Up @@ -227,7 +339,7 @@ mod tests {
"7.8.9",
);

reporter.send("list_assets").await.unwrap();
reporter.send("list_assets", None).await.unwrap();
let request = String::from_utf8(server.await.unwrap()).unwrap();
let (headers, _) = request.split_once("\r\n\r\n").unwrap();

Expand Down
5 changes: 4 additions & 1 deletion rust/crates/sift_mcp/src/server/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -95,7 +95,10 @@ impl ServerHandler for SiftMcpServer {
request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
let _ = self.client_event_reporter.send(request.name.as_ref()).await;
let _ = self
.client_event_reporter
.send(request.name.as_ref(), request.arguments.as_ref())
.await;

let context = ToolCallContext::new(self, request, context);
self.tool_router.call(context).await
Expand Down
Loading