Skip to main content

soma_mcp_client/upstream/
relay.rs

1use serde_json::Value;
2use thiserror::Error;
3
4pub mod cache;
5pub mod lifecycle;
6pub mod session;
7
8pub use cache::{RelayCache, RelayCacheKey, RelayConnectSlot, RelayConnection};
9pub use session::{RelaySessionId, RelaySessionMint};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
12pub enum RelayOperation {
13    CallTool,
14    ListTools,
15    ListResources,
16    GetPrompt,
17}
18
19#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
20pub struct RelayCapabilities {
21    pub elicitation: bool,
22    pub sampling: bool,
23    pub roots: bool,
24}
25
26impl RelayCapabilities {
27    #[must_use]
28    pub fn mirrored_by(self, downstream: Self) -> bool {
29        (!self.elicitation || downstream.elicitation)
30            && (!self.sampling || downstream.sampling)
31            && (!self.roots || downstream.roots)
32    }
33}
34
35#[derive(Debug, Error, PartialEq, Eq)]
36pub enum RelayError {
37    #[error("relay sessions only support call_tool")]
38    UnsupportedOperation,
39    #[error("relay session ids are gateway-minted and cannot come from user input")]
40    ForgedSessionId,
41    #[error("downstream client cannot mirror upstream relay capabilities")]
42    CapabilityMirrorMissing,
43}
44
45pub fn ensure_call_tool_only(operation: RelayOperation) -> Result<(), RelayError> {
46    if matches!(operation, RelayOperation::CallTool) {
47        return Ok(());
48    }
49    Err(RelayError::UnsupportedOperation)
50}
51
52pub fn reject_user_supplied_session_ids(params: &Value) -> Result<(), RelayError> {
53    let Some(object) = params.as_object() else {
54        return Ok(());
55    };
56    let forbidden = ["session_id", "mcp-session-id", "mcp_session_id"];
57    if forbidden.iter().any(|key| object.contains_key(*key)) {
58        return Err(RelayError::ForgedSessionId);
59    }
60    Ok(())
61}
62
63pub fn ensure_capabilities_mirrored(
64    upstream: RelayCapabilities,
65    downstream: RelayCapabilities,
66) -> Result<(), RelayError> {
67    if upstream.mirrored_by(downstream) {
68        return Ok(());
69    }
70    Err(RelayError::CapabilityMirrorMissing)
71}
72
73#[cfg(test)]
74#[path = "relay_tests.rs"]
75mod tests;