use std::sync::Arc;
use futures::Stream;
use objectiveai_sdk::agent::ClientObjectiveaiMcpEntry;
use objectiveai_sdk::cli::command::{AgentArguments, CommandExecutor, plugins, tools};
use rmcp::model::{
ClientCapabilities, ClientJsonRpcMessage, ClientRequest, GetExtensions, Implementation,
InitializeRequestParams, JsonRpcRequest, JsonRpcVersion2_0, NumberOrString, ProtocolVersion,
Request, ServerJsonRpcMessage,
};
use rmcp::service::serve_server;
use rmcp::transport::TransportAdapterIdentity;
use rmcp::transport::WorkerTransport;
use rmcp::transport::streamable_http_server::session::SessionManager;
use rmcp::transport::streamable_http_server::session::local::{
LocalSessionHandle, LocalSessionManager, LocalSessionManagerError, SessionConfig,
SessionError, create_local_session,
};
use rmcp::transport::streamable_http_server::session::{ServerSseMessage, SessionId};
use crate::agent_args_registry::{AgentArgumentsRegistry, SessionState};
use crate::objectiveai::ObjectiveAiMcpCli;
const HEADER_TO_FIELD: [(&str, fn(&mut AgentArguments, String)); 6] = [
(
"x-objectiveai-agent-instance-hierarchy",
|a, v| a.agent_instance_hierarchy = Some(v),
),
("x-objectiveai-agent-id", |a, v| a.agent_id = Some(v)),
("x-objectiveai-agent-full-id", |a, v| a.agent_full_id = Some(v)),
("x-objectiveai-agent-remote", |a, v| a.agent_remote = Some(v)),
("x-objectiveai-response-id", |a, v| a.response_id = Some(v)),
("x-objectiveai-response-ids", |a, v| a.response_ids = Some(v)),
];
#[derive(Debug, Clone)]
pub struct HeaderSessionManager<E> {
inner: Arc<LocalSessionManager>,
registry: Arc<AgentArgumentsRegistry>,
service: ObjectiveAiMcpCli<E>,
tools_list: Arc<Vec<tools::list::ResponseItem>>,
plugins_list: Arc<Vec<plugins::list::ResponseItem>>,
}
impl<E> HeaderSessionManager<E>
where
E: CommandExecutor + Send + Sync + 'static,
E::Error: std::fmt::Display + Send + 'static,
{
pub fn new(
registry: Arc<AgentArgumentsRegistry>,
service: ObjectiveAiMcpCli<E>,
tools_list: Arc<Vec<tools::list::ResponseItem>>,
plugins_list: Arc<Vec<plugins::list::ResponseItem>>,
) -> Self {
Self {
inner: Arc::new(LocalSessionManager::default()),
registry,
service,
tools_list,
plugins_list,
}
}
async fn mint_worker(
&self,
id: &SessionId,
message: &ClientJsonRpcMessage,
) -> Result<LocalSessionHandle, LocalSessionManagerError> {
let mut args = extract_agent_args(message);
args.mcp_session_id = Some(id.to_string());
let (mcp_root, mcp_tools, mcp_plugins) = extract_mcp_filter(message)?;
validate_mcp_filter(
mcp_tools.as_deref(),
mcp_plugins.as_deref(),
&self.tools_list,
&self.plugins_list,
)?;
let state = Arc::new(SessionState {
args,
mcp_root,
mcp_tools,
mcp_plugins,
});
self.registry.record(id.clone(), state).await;
let (handle, worker) = create_local_session(id.clone(), SessionConfig::default());
let transport = WorkerTransport::spawn(worker);
let svc = self.service.clone();
let id_for_close = id.clone();
let registry_for_close = self.registry.clone();
let inner_for_close = self.inner.clone();
tokio::spawn(async move {
let res =
serve_server::<_, _, _, TransportAdapterIdentity>(svc, transport).await;
if let Ok(svc) = res {
let _ = svc.waiting().await;
}
let _ = registry_for_close.remove(&id_for_close).await;
inner_for_close
.sessions
.write()
.await
.remove(&id_for_close);
});
Ok(handle)
}
async fn ensure_session(
&self,
id: &SessionId,
message: &ClientJsonRpcMessage,
) -> Result<(), LocalSessionManagerError> {
if self.inner.has_session(id).await? {
return Ok(());
}
let handle = self.mint_worker(id, message).await?;
handle
.initialize(synthetic_initialize_message())
.await
.map_err(|e| error_invalid_input(format!("synthetic initialize: {e}")))?;
self.inner.sessions.write().await.insert(id.clone(), handle);
Ok(())
}
}
impl<E> SessionManager for HeaderSessionManager<E>
where
E: CommandExecutor + Send + Sync + 'static,
E::Error: std::fmt::Display + Send + 'static,
{
type Error = LocalSessionManagerError;
type Transport = <LocalSessionManager as SessionManager>::Transport;
async fn create_session(&self) -> Result<(SessionId, Self::Transport), Self::Error> {
self.inner.create_session().await
}
async fn initialize_session(
&self,
id: &SessionId,
message: ClientJsonRpcMessage,
) -> Result<ServerJsonRpcMessage, Self::Error> {
let mut args = extract_agent_args(&message);
args.mcp_session_id = Some(id.to_string());
let (mcp_root, mcp_tools, mcp_plugins) = extract_mcp_filter(&message)?;
validate_mcp_filter(
mcp_tools.as_deref(),
mcp_plugins.as_deref(),
&self.tools_list,
&self.plugins_list,
)?;
let state = Arc::new(SessionState {
args,
mcp_root,
mcp_tools,
mcp_plugins,
});
self.registry.record(id.clone(), state).await;
self.inner.initialize_session(id, message).await
}
async fn has_session(&self, _id: &SessionId) -> Result<bool, Self::Error> {
Ok(true)
}
async fn close_session(&self, id: &SessionId) -> Result<(), Self::Error> {
let _ = self.registry.remove(id).await;
self.inner.close_session(id).await
}
async fn create_stream(
&self,
id: &SessionId,
message: ClientJsonRpcMessage,
) -> Result<impl Stream<Item = ServerSseMessage> + Send + Sync + 'static, Self::Error> {
if is_initialize(&message) && !self.inner.has_session(id).await? {
let handle = self.mint_worker(id, &message).await?;
let response = handle
.initialize(message)
.await
.map_err(|e| error_invalid_input(format!("resume initialize: {e}")))?;
self.inner.sessions.write().await.insert(id.clone(), handle);
let item = ServerSseMessage::from_message(response);
let stream: std::pin::Pin<
Box<dyn Stream<Item = ServerSseMessage> + Send + Sync + 'static>,
> = Box::pin(futures::stream::iter(vec![item]));
return Ok(stream);
}
self.ensure_session(id, &message).await?;
let inner = self.inner.create_stream(id, message).await?;
let stream: std::pin::Pin<
Box<dyn Stream<Item = ServerSseMessage> + Send + Sync + 'static>,
> = Box::pin(inner);
Ok(stream)
}
async fn accept_message(
&self,
id: &SessionId,
message: ClientJsonRpcMessage,
) -> Result<(), Self::Error> {
self.ensure_session(id, &message).await?;
self.inner.accept_message(id, message).await
}
async fn create_standalone_stream(
&self,
id: &SessionId,
) -> Result<impl Stream<Item = ServerSseMessage> + Send + Sync + 'static, Self::Error> {
self.inner.create_standalone_stream(id).await
}
async fn resume(
&self,
id: &SessionId,
last_event_id: String,
) -> Result<impl Stream<Item = ServerSseMessage> + Send + Sync + 'static, Self::Error> {
self.inner.resume(id, last_event_id).await
}
}
fn is_initialize(m: &ClientJsonRpcMessage) -> bool {
matches!(
m,
ClientJsonRpcMessage::Request(r)
if matches!(r.request, ClientRequest::InitializeRequest(_))
)
}
fn extract_agent_args(message: &ClientJsonRpcMessage) -> AgentArguments {
let parts = match message {
ClientJsonRpcMessage::Request(r) => {
r.request.extensions().get::<http::request::Parts>()
}
ClientJsonRpcMessage::Notification(n) => {
n.notification.extensions().get::<http::request::Parts>()
}
_ => None,
};
let mut args = AgentArguments::default();
if let Some(p) = parts {
for (name, setter) in HEADER_TO_FIELD {
if let Some(v) = p.headers.get(name).and_then(|v| v.to_str().ok()) {
let s = v.trim();
if !s.is_empty() {
setter(&mut args, s.to_string());
}
}
}
}
args
}
fn synthetic_initialize_message() -> ClientJsonRpcMessage {
let mut client_info = Implementation::default();
client_info.name = "objectiveai-mcp-restore-stub".into();
client_info.version = "0".into();
let params = InitializeRequestParams::new(ClientCapabilities::default(), client_info)
.with_protocol_version(ProtocolVersion::V_2025_06_18);
let request = Request::new(params);
ClientJsonRpcMessage::Request(JsonRpcRequest {
jsonrpc: JsonRpcVersion2_0,
id: NumberOrString::Number(0),
request: ClientRequest::InitializeRequest(request),
})
}
fn error_invalid_input(msg: String) -> LocalSessionManagerError {
LocalSessionManagerError::SessionError(SessionError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
msg,
)))
}
fn extract_mcp_filter(
message: &ClientJsonRpcMessage,
) -> Result<
(
bool,
Option<Vec<ClientObjectiveaiMcpEntry>>,
Option<Vec<ClientObjectiveaiMcpEntry>>,
),
LocalSessionManagerError,
> {
let parts = match message {
ClientJsonRpcMessage::Request(r) => {
r.request.extensions().get::<http::request::Parts>()
}
ClientJsonRpcMessage::Notification(n) => {
n.notification.extensions().get::<http::request::Parts>()
}
_ => None,
};
let Some(p) = parts else {
return Ok((true, None, None));
};
let mut root = true;
let mut tools: Option<Vec<ClientObjectiveaiMcpEntry>> = None;
let mut plugins: Option<Vec<ClientObjectiveaiMcpEntry>> = None;
if let Some(v) = p.headers.get("x-objectiveai-mcp-root").and_then(|v| v.to_str().ok()) {
let s = v.trim();
root = match s {
"true" => true,
"false" => false,
other => {
return Err(error_invalid_input(format!(
"x-objectiveai-mcp-root must be \"true\" or \"false\", got {other:?}"
)));
}
};
}
if let Some(v) = p.headers.get("x-objectiveai-mcp-tools").and_then(|v| v.to_str().ok()) {
let s = v.trim();
if !s.is_empty() {
let parsed: Vec<ClientObjectiveaiMcpEntry> =
serde_json::from_str(s).map_err(|e| {
error_invalid_input(format!(
"x-objectiveai-mcp-tools: invalid JSON ({e})"
))
})?;
tools = Some(parsed);
}
}
if let Some(v) = p.headers.get("x-objectiveai-mcp-plugins").and_then(|v| v.to_str().ok()) {
let s = v.trim();
if !s.is_empty() {
let parsed: Vec<ClientObjectiveaiMcpEntry> =
serde_json::from_str(s).map_err(|e| {
error_invalid_input(format!(
"x-objectiveai-mcp-plugins: invalid JSON ({e})"
))
})?;
plugins = Some(parsed);
}
}
Ok((root, tools, plugins))
}
fn validate_mcp_filter(
tools: Option<&[ClientObjectiveaiMcpEntry]>,
plugins: Option<&[ClientObjectiveaiMcpEntry]>,
tools_list: &[tools::list::ResponseItem],
plugins_list: &[plugins::list::ResponseItem],
) -> Result<(), LocalSessionManagerError> {
if let Some(declared) = tools {
for entry in declared {
let found = tools_list.iter().any(|t| {
t.owner == entry.owner
&& t.name == entry.name
&& t.version == entry.version
});
if !found {
return Err(error_invalid_input(format!(
"tool {}/{}@{} not installed",
entry.owner, entry.name, entry.version
)));
}
}
}
if let Some(declared) = plugins {
for entry in declared {
let found = plugins_list.iter().any(|p| {
p.owner == entry.owner
&& p.name == entry.name
&& p.version == entry.version
});
if !found {
return Err(error_invalid_input(format!(
"plugin {}/{}@{} not installed",
entry.owner, entry.name, entry.version
)));
}
}
}
Ok(())
}