use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use futures::StreamExt;
use futures::stream::BoxStream;
use rmcp::transport::streamable_http_client::{
StreamableHttpClient, StreamableHttpError, StreamableHttpPostResponse,
};
use serde_json::Value;
use sse_stream::{Error as SseError, Sse};
use crate::http::HttpTransport;
#[derive(Clone)]
pub struct AgentdHttp {
http: Arc<HttpTransport>,
timeout: Duration,
}
impl AgentdHttp {
pub fn new(http: Arc<HttpTransport>, timeout: Duration) -> AgentdHttp {
AgentdHttp { http, timeout }
}
}
#[derive(Debug, thiserror::Error)]
#[error("{0}")]
pub struct TransportError(String);
enum Pumped {
Message(Value),
Done(Result<Option<Value>, String>),
}
fn as_event(v: &Value) -> Sse {
Sse {
event: None,
data: Some(v.to_string()),
id: None,
retry: None,
}
}
impl StreamableHttpClient for AgentdHttp {
type Error = TransportError;
async fn post_message(
&self,
_uri: Arc<str>,
message: rmcp::model::ClientJsonRpcMessage,
_session_id: Option<Arc<str>>,
auth_header: Option<String>,
custom_headers: HashMap<http::HeaderName, http::HeaderValue>,
) -> Result<StreamableHttpPostResponse, StreamableHttpError<Self::Error>> {
let body = serde_json::to_vec(&message)
.map_err(|e| StreamableHttpError::Client(TransportError(e.to_string())))?;
let request_id = request_id_of(&message);
let http = Arc::clone(&self.http);
let timeout = self.timeout;
let extra = header_pairs(auth_header, custom_headers);
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<Pumped>();
let notes_tx = tx.clone();
tokio::task::spawn_blocking(move || {
let refs: Vec<(&str, &str)> = extra
.iter()
.map(|(k, v)| (k.as_str(), v.as_str()))
.collect();
let resp = http.send(request_id, &body, timeout, &refs, |n| {
let _ = notes_tx.send(Pumped::Message(n));
});
let _ = tx.send(Pumped::Done(resp.map_err(|e| e.to_string())));
});
let first = match rx.recv().await {
Some(Pumped::Message(v)) | Some(Pumped::Done(Ok(Some(v)))) => v,
Some(Pumped::Done(Ok(None))) => return Ok(StreamableHttpPostResponse::Accepted),
Some(Pumped::Done(Err(e))) => {
return Err(StreamableHttpError::Client(TransportError(e)));
}
None => {
return Err(StreamableHttpError::Client(TransportError(
"mcp: response stream ended with no reply".into(),
)));
}
};
let session = self.http.session_id();
let rest = futures::stream::unfold(rx, |mut rx| async move {
match rx.recv().await {
Some(Pumped::Message(v)) | Some(Pumped::Done(Ok(Some(v)))) => {
Some((Ok(as_event(&v)), rx))
}
_ => None,
}
});
let head: Vec<Result<Sse, SseError>> = vec![Ok(as_event(&first))];
Ok(StreamableHttpPostResponse::Sse(
Box::pin(futures::stream::iter(head).chain(rest)),
session,
))
}
async fn delete_session(
&self,
_uri: Arc<str>,
session_id: Arc<str>,
auth_header: Option<String>,
custom_headers: HashMap<http::HeaderName, http::HeaderValue>,
) -> Result<(), StreamableHttpError<Self::Error>> {
let http = Arc::clone(&self.http);
let timeout = self.timeout;
let extra = header_pairs(auth_header, custom_headers);
let sid = session_id.to_string();
let _ = (http, timeout, extra, sid);
Ok(())
}
async fn get_stream(
&self,
_uri: Arc<str>,
_session_id: Option<Arc<str>>,
last_event_id: Option<String>,
auth_header: Option<String>,
custom_headers: HashMap<http::HeaderName, http::HeaderValue>,
) -> Result<BoxStream<'static, Result<Sse, SseError>>, StreamableHttpError<Self::Error>> {
let http = Arc::clone(&self.http);
let timeout = self.timeout;
let extra = header_pairs(auth_header, custom_headers);
let (tx, rx) = tokio::sync::mpsc::unbounded_channel::<Result<Sse, SseError>>();
std::thread::spawn(move || {
let mut refs: Vec<(&str, &str)> = extra
.iter()
.map(|(k, v)| (k.as_str(), v.as_str()))
.collect();
if let Some(id) = &last_event_id {
refs.push(("Last-Event-ID", id.as_str()));
}
let _ = &refs;
let Ok(mut events) = http.open_events(timeout) else {
return;
};
while let Ok(Some(ev)) = events.next_event() {
let sse = Sse {
event: ev.event,
data: Some(ev.data),
id: ev.id,
retry: None,
};
if tx.send(Ok(sse)).is_err() {
return;
}
}
});
Ok(Box::pin(
tokio_stream::wrappers::UnboundedReceiverStream::new(rx),
))
}
}
fn request_id_of(message: &rmcp::model::ClientJsonRpcMessage) -> Option<i64> {
serde_json::to_value(message)
.ok()
.and_then(|v| v.get("id").and_then(Value::as_i64))
}
fn header_pairs(
auth_header: Option<String>,
custom: HashMap<http::HeaderName, http::HeaderValue>,
) -> Vec<(String, String)> {
let mut out = Vec::new();
if let Some(a) = auth_header {
out.push(("Authorization".to_string(), a));
}
for (k, v) in custom {
if let Ok(s) = v.to_str() {
out.push((k.as_str().to_string(), s.to_string()));
}
}
out
}