use rmcp::model::CallToolRequestParam;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use exocortex_wire::ingest::v1::{
ingest_service_client::IngestServiceClient, IngestBatch, MemoryDraft as WireMemoryDraft,
ProducerIdentity, RegisterSourceRequest,
};
#[derive(Debug, Clone, Deserialize, Serialize, schemars::JsonSchema)]
pub struct EndSessionArgs {
#[serde(default)]
pub session_id: Option<String>,
pub project_id: String,
#[serde(default)]
pub team_id: Option<String>,
pub memories: Vec<MemoryDraftInput>,
#[serde(default)]
pub edges: Vec<EdgeHintInput>,
}
pub type MemoryDraftInput = crate::preflight::PreflightMemoryDraft;
pub type EdgeHintInput = crate::preflight::PreflightEdgeHint;
#[derive(Debug, Clone, Serialize, schemars::JsonSchema)]
pub struct EndSessionAck {
pub accepted: u32,
pub rejected: u32,
pub assigned_lsn: u64,
pub rejections: Vec<RejectionSummary>,
pub local_validation_failed: bool,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub unverified: Vec<crate::preflight::UnverifiedCheck>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub similar_to: Vec<SimilarToSummary>,
}
#[derive(Debug, Clone, Serialize, schemars::JsonSchema)]
pub struct RejectionSummary {
pub draft_key: String,
pub code: String,
pub detail: String,
#[serde(default, skip_serializing_if = "String::is_empty")]
pub correction: String,
}
#[derive(Debug, Clone, Serialize, schemars::JsonSchema)]
pub struct SimilarToSummary {
pub draft_key: String,
pub existing_memory_id: String,
pub existing_title: String,
pub suggestion: String,
}
pub struct EndSessionTool {
pub client: IngestServiceClient<tonic::transport::Channel>,
pub org_id: String,
pub fingerprint: [u8; 32],
pub hmac_key: [u8; 32],
pub node_id: String,
pub agent_id: String,
pub auth_token: Option<String>,
pub ontology: Arc<exocortex_kernel::Ontology>,
pub cache: Option<Arc<exocortex_cache::LocalCache>>,
pub vc: exocortex_ops::VisibilityContext,
}
fn parse_visibility(s: &str) -> Result<i32, String> {
match s.to_lowercase().as_str() {
"private" => Ok(0),
"project" => Ok(1),
"team" => Ok(2),
"org" => Ok(3),
other => Err(format!("unknown visibility `{other}`")),
}
}
impl EndSessionTool {
pub async fn handle(&self, args: EndSessionArgs) -> Result<EndSessionAck, rmcp::Error> {
let session_id = args.session_id.clone().unwrap_or_default();
if session_id.is_empty() {
return Err(rmcp::Error::invalid_params(
"session_id: the MCP layer must stamp the client-minted conversation id (§4.8)",
None,
));
}
if args.edges.len() > exocortex_wire::limits::MAX_EDGES_PER_BATCH {
return Err(rmcp::Error::invalid_params(
"edges: at most 64 relationships per request",
None,
));
}
for memory in &args.memories {
if let Err(detail) =
exocortex_wire::limits::validate_memory_fields(&memory.content, &memory.tags)
{
return Err(rmcp::Error::invalid_params(detail, None));
}
}
let cache = self.cache.clone();
let org = self.org_id.clone();
let mut vc = self.vc.clone();
if !args.project_id.is_empty()
&& !vc
.project_ids
.iter()
.any(|id| id.as_str() == args.project_id)
{
vc.project_ids.push(args.project_id.clone().into());
}
if let Some(team_id) = args.team_id.as_deref().filter(|id| !id.is_empty()) {
if !vc.team_ids.iter().any(|id| id.as_str() == team_id) {
vc.team_ids.push(team_id.into());
}
}
let pre =
crate::preflight::validate_batch(&self.ontology, &args.memories, &args.edges, |id| {
cache.as_ref().and_then(|c| {
let id = exocortex_kernel::MemoryId::parse_hex(id)?;
c.get_memory(&org, &id, &vc).map(|m| m.memory_type)
})
});
if !pre.rejections.is_empty() {
return Ok(EndSessionAck {
accepted: 0,
rejected: pre.would_reject,
assigned_lsn: 0,
rejections: pre
.rejections
.iter()
.map(|r| RejectionSummary {
draft_key: r.draft_key.clone(),
code: r.code.clone(),
detail: r.detail.clone(),
correction: r.correction.clone(),
})
.collect(),
local_validation_failed: true,
unverified: pre.unverified,
similar_to: vec![],
});
}
let unverified = pre.unverified;
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default();
let ts = prost_types::Timestamp {
seconds: now.as_secs() as i64,
nanos: now.subsec_nanos() as i32,
};
let draft_vis: std::collections::HashMap<String, i32> = args
.memories
.iter()
.map(|m| {
let v = match m.visibility.to_lowercase().as_str() {
"private" => 0,
"project" => 1,
"team" => 2,
_ => 3,
};
(m.draft_key.clone(), v)
})
.collect();
let mut memories = Vec::with_capacity(args.memories.len());
for m in args.memories {
let vis = parse_visibility(&m.visibility)
.map_err(|_| rmcp::Error::invalid_params("unknown visibility", None))?;
memories.push(WireMemoryDraft {
rights: None,
draft_key: m.draft_key,
id: uuid::Uuid::now_v7().simple().to_string(),
memory_type: m.memory_type,
title: m.title,
content: m.content,
tags: m.tags,
visibility: vis,
valid_from: Some(ts),
valid_until: None,
external_key: None, });
}
let relationships: Vec<exocortex_wire::ingest::v1::RelationshipDraft> = args
.edges
.into_iter()
.map(|e| exocortex_wire::ingest::v1::RelationshipDraft {
from_draft_key: e.from_draft_key.clone(),
to_draft_key: e.to_draft_key.clone(),
kind: e.kind,
strength: e.strength,
confidence: 0.8,
context: String::new(),
visibility: if e.to_draft_key.is_empty() {
draft_vis
.get(&e.from_draft_key)
.copied()
.unwrap_or(1)
.min(3)
} else {
draft_vis
.get(&e.from_draft_key)
.copied()
.unwrap_or(1)
.min(draft_vis.get(&e.to_draft_key).copied().unwrap_or(1))
},
to_memory_id: e.to_memory_id,
})
.collect();
let mut batch = IngestBatch {
org_id: self.org_id.clone(),
source_uri: format!("session://{session_id}"),
producer_id: "session-wrapup".into(),
batch_id: crate::drain::content_batch_id(&session_id, &memories, &relationships),
mapping_version: "session-wrapup:1.0.0".into(),
ontology_fingerprint: self.fingerprint.to_vec(),
ceiling: 3, checksum: String::new(), observed_at: Some(ts),
recorded_at: Some(ts),
snapshot: None, memories,
relationships,
producer: Some(ProducerIdentity {
node_id: self.node_id.clone(),
agent_id: self.agent_id.clone(),
adapter_id: String::new(),
hmac_signature: vec![],
client_metadata: Some(exocortex_wire::ingest::v1::ClientMetadata {
playbook_version: crate::playbook::PLAYBOOK_VERSION.into(),
client_version: env!("CARGO_PKG_VERSION").into(),
harness_hint: String::new(),
project_id: args.project_id,
team_id: args.team_id.unwrap_or_default(),
}),
}),
};
exocortex_wire::signing::prepare_batch(&self.hmac_key, &mut batch);
let mut client = self.client.clone();
let mut registration = RegisterSourceRequest {
default_rights: None,
org_id: self.org_id.clone(),
source_uri: batch.source_uri.clone(),
producer_id: "session-wrapup".into(),
ceiling: 3,
source_flavor: "session".into(),
projection: None,
producer_kind: exocortex_wire::ingest::v1::ProducerKind::CodingAgent.into(),
producer: Some(ProducerIdentity {
node_id: self.node_id.clone(),
agent_id: self.agent_id.clone(),
adapter_id: String::new(),
hmac_signature: vec![],
client_metadata: batch
.producer
.as_ref()
.and_then(|p| p.client_metadata.clone()),
}),
};
exocortex_wire::signing::sign_registration(&self.hmac_key, &mut registration);
let mut reg_req = tonic::Request::new(registration);
reg_req.set_timeout(std::time::Duration::from_secs(20));
if let Some(token) = &self.auth_token {
if let Ok(v) = format!("Bearer {token}").parse() {
reg_req.metadata_mut().insert("authorization", v);
}
}
if let Err(e) = client.register_source(reg_req).await {
tracing::warn!(%e, "register_source failed; submit will surface the cause");
}
let mut submit_req = tonic::Request::new(batch);
submit_req.set_timeout(std::time::Duration::from_secs(30));
if let Some(token) = &self.auth_token {
if let Ok(v) = format!("Bearer {token}").parse() {
submit_req.metadata_mut().insert("authorization", v);
}
}
let ack = client
.submit(submit_req)
.await
.map_err(|e| rmcp::Error::internal_error(format!("ingest: {e}"), None))?
.into_inner();
Ok(EndSessionAck {
accepted: ack.accepted,
rejected: ack.rejected,
assigned_lsn: ack.assigned_lsn,
rejections: ack
.rejections
.iter()
.map(|r| {
let code = exocortex_wire::ingest::v1::RejectCode::try_from(r.code)
.unwrap_or(exocortex_wire::ingest::v1::RejectCode::Unknown);
RejectionSummary {
draft_key: r.draft_key.clone(),
code: format!("{code:?}"),
detail: r.detail.clone(),
correction: exocortex_wire::corrections::guidance(code)
.correction
.into(),
}
})
.collect(),
local_validation_failed: false,
unverified,
similar_to: ack
.similar_to
.iter()
.map(|h| SimilarToSummary {
draft_key: h.draft_key.clone(),
existing_memory_id: h.existing_memory_id.clone(),
existing_title: h.existing_title.clone(),
suggestion: h.suggestion.clone(),
})
.collect(),
})
}
}
pub fn parse_args(request: CallToolRequestParam) -> Result<EndSessionArgs, rmcp::Error> {
let tcc = ToolCallContextArgs { request };
serde_json::from_value(serde_json::Value::Object(tcc.arguments()))
.map_err(|_| rmcp::Error::invalid_params("bad end_session args", None))
}
struct ToolCallContextArgs {
request: CallToolRequestParam,
}
impl ToolCallContextArgs {
fn arguments(&self) -> serde_json::Map<String, serde_json::Value> {
self.request.arguments.clone().unwrap_or_default()
}
}