soma_mcp_client/upstream/
relay.rs1use 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;