use serde::{Deserialize, Serialize};
use crate::mcp_server::ScopedMcpServer;
use crate::session::ExecutionSession;
use crate::session_services::SessionStorageStore;
use crate::typed_id::SessionId;
pub const SESSION_MCP_SERVER_KV_PREFIX: &str = "session_mcp:";
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum SessionMcpServerSource {
Ard {
urn: String,
},
UserMcp,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct SessionMcpServer {
pub name: String,
pub server: ScopedMcpServer,
pub source: SessionMcpServerSource,
}
pub fn session_mcp_server_kv_key(name: &str) -> String {
format!("{SESSION_MCP_SERVER_KV_PREFIX}{name}")
}
pub async fn put_session_mcp_server(
storage: &dyn SessionStorageStore,
session_id: SessionId,
record: &SessionMcpServer,
) -> everruns_contracts::error::Result<()> {
let serialized = serde_json::to_string(record)
.map_err(|e| everruns_contracts::error::AgentLoopError::Internal(e.into()))?;
storage
.set_value(
session_id,
&session_mcp_server_kv_key(&record.name),
&serialized,
)
.await
}
pub async fn get_session_mcp_server(
storage: &dyn SessionStorageStore,
session_id: SessionId,
name: &str,
) -> Option<SessionMcpServer> {
let raw = storage
.get_value(session_id, &session_mcp_server_kv_key(name))
.await
.ok()??;
serde_json::from_str(&raw).ok()
}
pub async fn remove_session_mcp_server(
storage: &dyn SessionStorageStore,
session_id: SessionId,
name: &str,
) -> everruns_contracts::error::Result<bool> {
storage
.delete_value(session_id, &session_mcp_server_kv_key(name))
.await
}
pub async fn load_session_mcp_servers(
storage: &dyn SessionStorageStore,
session_id: SessionId,
) -> Vec<SessionMcpServer> {
let keys = match storage.list_keys(session_id).await {
Ok(keys) => keys,
Err(e) => {
tracing::warn!("failed to list session keys for session MCP servers: {e}");
return Vec::new();
}
};
let mut records = Vec::new();
for info in keys {
if !info.key.starts_with(SESSION_MCP_SERVER_KV_PREFIX) {
continue;
}
match storage.get_value(session_id, &info.key).await {
Ok(Some(raw)) => match serde_json::from_str::<SessionMcpServer>(&raw) {
Ok(record) => records.push(record),
Err(e) => {
tracing::warn!(key = %info.key, "skipping malformed session MCP server: {e}")
}
},
Ok(None) => {}
Err(e) => tracing::warn!(key = %info.key, "failed to read session MCP server: {e}"),
}
}
records
}
pub fn merge_session_mcp_servers(session: &mut ExecutionSession, records: &[SessionMcpServer]) {
for record in records {
session
.mcp_servers
.insert(record.name.clone(), record.server.clone());
}
}
#[cfg(test)]
#[path = "session_mcp_servers_tests.rs"]
mod tests;