Skip to main content

systemprompt_api/services/proxy/engine/
mcp_session.rs

1//! Proxy-side identity store for MCP sessions.
2//!
3//! MCP clients authenticate on the `initialize` call but may omit the bearer
4//! token on subsequent session-only requests. This module persists the
5//! authenticated identity keyed by `mcp-session-id` so those follow-ups can be
6//! enriched ([`enrich_with_cached_identity`]) on any replica, and evicts the
7//! row on session teardown or a stale-session backend response
8//! ([`handle_mcp_response`]). The store is the trust anchor for session-based
9//! MCP auth — rows are only written for a verified [`AuthenticatedUser`].
10//!
11//! Copyright (c) systemprompt.io — Business Source License 1.1.
12//! See <https://systemprompt.io> for licensing details.
13
14use axum::http::{HeaderMap, StatusCode};
15use systemprompt_identifiers::{SessionId, UserId};
16use systemprompt_mcp::repository::{McpProxyIdentityRepository, ProxyIdentityRow};
17use systemprompt_models::RequestContext;
18use systemprompt_models::auth::AuthenticatedUser;
19use uuid::Uuid;
20
21fn session_id_header(headers: &HeaderMap) -> Option<SessionId> {
22    headers
23        .get("mcp-session-id")
24        .and_then(|v| v.to_str().ok())
25        .map(|s| SessionId::new(s.to_owned()))
26}
27
28async fn evict(identities: &McpProxyIdentityRepository, session_id: &SessionId, reason: &str) {
29    if let Err(e) = identities.delete(session_id).await {
30        tracing::warn!(session_id = %session_id, error = %e, reason, "Failed to evict proxy session identity");
31    }
32}
33
34pub async fn enrich_with_cached_identity(
35    identities: &McpProxyIdentityRepository,
36    request_headers: &HeaderMap,
37    req_context: RequestContext,
38    service_name: &str,
39) -> RequestContext {
40    let Some(session_id) = session_id_header(request_headers) else {
41        return req_context;
42    };
43
44    let identity = match identities.find(&session_id).await {
45        Ok(Some(identity)) => identity,
46        Ok(None) => {
47            tracing::debug!(
48                service = %service_name,
49                session_id = %session_id,
50                "No stored identity for session-only MCP request"
51            );
52            return req_context;
53        },
54        Err(e) => {
55            tracing::warn!(
56                service = %service_name,
57                session_id = %session_id,
58                error = %e,
59                "Proxy session identity lookup failed"
60            );
61            return req_context;
62        },
63    };
64
65    let Ok(user_uuid) = Uuid::parse_str(identity.user_id.as_str()) else {
66        tracing::warn!(
67            service = %service_name,
68            session_id = %session_id,
69            user_id = %identity.user_id,
70            "Stored proxy session identity has a non-UUID user id"
71        );
72        return req_context;
73    };
74
75    tracing::info!(
76        service = %service_name,
77        session_id = %session_id,
78        user_id = %identity.user_id,
79        "Enriching session-only request with stored identity"
80    );
81    req_context
82        .with_actor(systemprompt_identifiers::Actor::user(identity.user_id))
83        .with_user_type(identity.user_type)
84        .with_auth_token(identity.auth_token.as_str().to_owned())
85        .with_user(AuthenticatedUser::new(
86            user_uuid,
87            String::new(),
88            String::new(),
89            identity.permissions,
90        ))
91}
92
93#[derive(Debug)]
94pub struct McpResponseCtx<'a> {
95    pub identities: &'a McpProxyIdentityRepository,
96    pub response: &'a reqwest::Response,
97    pub request_headers: &'a HeaderMap,
98    pub req_context: &'a RequestContext,
99    pub authenticated_user: Option<&'a AuthenticatedUser>,
100    pub service_name: &'a str,
101    pub method_str: &'a str,
102}
103
104pub async fn handle_mcp_response(args: McpResponseCtx<'_>) {
105    let McpResponseCtx {
106        identities,
107        response,
108        request_headers,
109        req_context,
110        authenticated_user,
111        service_name,
112        method_str,
113    } = args;
114    let resp_status = response.status();
115    let resp_session = response
116        .headers()
117        .get("mcp-session-id")
118        .and_then(|v| v.to_str().ok())
119        .unwrap_or("none");
120    let resp_content_type = response
121        .headers()
122        .get("content-type")
123        .and_then(|v| v.to_str().ok())
124        .unwrap_or("none");
125
126    tracing::info!(
127        service = %service_name,
128        status = %resp_status,
129        resp_session_id = %resp_session,
130        content_type = %resp_content_type,
131        method = %method_str,
132        "MCP backend response"
133    );
134
135    if !resp_status.is_success() {
136        evict_on_error_response(
137            identities,
138            response,
139            request_headers,
140            service_name,
141            method_str,
142        )
143        .await;
144    }
145
146    cache_identity_from_response(
147        identities,
148        response,
149        req_context,
150        authenticated_user,
151        service_name,
152    )
153    .await;
154
155    if method_str == "DELETE"
156        && let Some(session_id) = session_id_header(request_headers)
157    {
158        evict(identities, &session_id, "delete").await;
159        tracing::debug!(session_id = %session_id, "Evicted session identity on DELETE");
160    }
161}
162
163async fn evict_on_error_response(
164    identities: &McpProxyIdentityRepository,
165    response: &reqwest::Response,
166    request_headers: &HeaderMap,
167    service_name: &str,
168    method_str: &str,
169) {
170    let resp_status = response.status();
171    let header_dump: Vec<String> = response
172        .headers()
173        .iter()
174        .map(|(k, v)| format!("{}: {}", k, v.to_str().unwrap_or("?")))
175        .collect();
176    tracing::error!(
177        service = %service_name,
178        status = %resp_status,
179        headers = ?header_dump,
180        "MCP backend error response"
181    );
182
183    if resp_status == StatusCode::NOT_FOUND
184        && (method_str == "GET" || method_str == "POST")
185        && let Some(session_id) = session_id_header(request_headers)
186    {
187        evict(identities, &session_id, "stale_session").await;
188        tracing::info!(
189            service = %service_name,
190            session_id = %session_id,
191            method = %method_str,
192            reason = "stale_session",
193            "Evicted stale proxy session identity on 404; the backend no longer knows this session (restarted child?) and the client must re-initialise"
194        );
195    }
196}
197
198async fn cache_identity_from_response(
199    identities: &McpProxyIdentityRepository,
200    response: &reqwest::Response,
201    req_context: &RequestContext,
202    authenticated_user: Option<&AuthenticatedUser>,
203    service_name: &str,
204) {
205    let Some(session_id) = session_id_header(response.headers()) else {
206        return;
207    };
208    let Some(user) = authenticated_user else {
209        return;
210    };
211    let row = ProxyIdentityRow {
212        user_id: UserId::new(user.id.to_string()),
213        user_type: req_context.user_type(),
214        permissions: user.permissions.clone(),
215        auth_token: req_context.auth_token().clone(),
216    };
217    match identities.upsert(&session_id, &row).await {
218        Ok(()) => tracing::info!(
219            service = %service_name,
220            session_id = %session_id,
221            user_id = %user.id,
222            "Stored session identity for MCP session"
223        ),
224        Err(e) => tracing::error!(
225            service = %service_name,
226            session_id = %session_id,
227            user_id = %user.id,
228            error = %e,
229            "Failed to store session identity for MCP session"
230        ),
231    }
232}