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::{McpServerId, ServiceName, SessionId};
16use systemprompt_mcp::repository::{McpProxyIdentityRepository, ProxyIdentityRow};
17use systemprompt_models::RequestContext;
18use systemprompt_models::auth::AuthenticatedUser;
19
20fn session_id_header(headers: &HeaderMap) -> Option<SessionId> {
21    headers
22        .get("mcp-session-id")
23        .and_then(|v| v.to_str().ok())
24        .map(|s| SessionId::new(s.to_owned()))
25}
26
27async fn evict(identities: &McpProxyIdentityRepository, session_id: &SessionId, reason: &str) {
28    if let Err(e) = identities.delete(session_id).await {
29        tracing::warn!(session_id = %session_id, error = %e, reason, "Failed to evict proxy session identity");
30    }
31}
32
33pub async fn enrich_with_cached_identity(
34    identities: &McpProxyIdentityRepository,
35    request_headers: &HeaderMap,
36    req_context: RequestContext,
37    service_name: &ServiceName,
38) -> RequestContext {
39    let Some(session_id) = session_id_header(request_headers) else {
40        return req_context;
41    };
42
43    let identity = match identities.find(&session_id).await {
44        Ok(Some(identity)) => identity,
45        Ok(None) => {
46            tracing::debug!(
47                service = %service_name,
48                session_id = %session_id,
49                "No stored identity for session-only MCP request"
50            );
51            return req_context;
52        },
53        Err(e) => {
54            tracing::warn!(
55                service = %service_name,
56                session_id = %session_id,
57                error = %e,
58                "Proxy session identity lookup failed"
59            );
60            return req_context;
61        },
62    };
63
64    tracing::info!(
65        service = %service_name,
66        session_id = %session_id,
67        user_id = %identity.user_id,
68        "Enriching session-only request with stored identity"
69    );
70    req_context
71        .with_actor(systemprompt_identifiers::Actor::user(
72            identity.user_id.clone(),
73        ))
74        .with_user_type(identity.user_type)
75        .with_auth_token(identity.auth_token)
76        .with_user(AuthenticatedUser::new_with_roles(
77            identity.user_id,
78            String::new(),
79            String::new(),
80            identity.permissions,
81            identity.roles,
82        ))
83}
84
85#[derive(Debug)]
86pub struct McpResponseCtx<'a> {
87    pub identities: &'a McpProxyIdentityRepository,
88    pub response: &'a reqwest::Response,
89    pub request_headers: &'a HeaderMap,
90    pub req_context: &'a RequestContext,
91    pub authenticated_user: Option<&'a AuthenticatedUser>,
92    pub service_name: &'a ServiceName,
93    pub method_str: &'a str,
94}
95
96pub async fn handle_mcp_response(args: McpResponseCtx<'_>) {
97    let McpResponseCtx {
98        identities,
99        response,
100        request_headers,
101        req_context,
102        authenticated_user,
103        service_name,
104        method_str,
105    } = args;
106    let resp_status = response.status();
107    let resp_session = response
108        .headers()
109        .get("mcp-session-id")
110        .and_then(|v| v.to_str().ok())
111        .unwrap_or("none");
112    let resp_content_type = response
113        .headers()
114        .get("content-type")
115        .and_then(|v| v.to_str().ok())
116        .unwrap_or("none");
117
118    tracing::info!(
119        service = %service_name,
120        status = %resp_status,
121        resp_session_id = %resp_session,
122        content_type = %resp_content_type,
123        method = %method_str,
124        "MCP backend response"
125    );
126
127    if !resp_status.is_success() {
128        evict_on_error_response(
129            identities,
130            response,
131            request_headers,
132            service_name,
133            method_str,
134        )
135        .await;
136    }
137
138    cache_identity_from_response(
139        identities,
140        response,
141        req_context,
142        authenticated_user,
143        service_name,
144    )
145    .await;
146
147    if method_str == "DELETE"
148        && let Some(session_id) = session_id_header(request_headers)
149    {
150        evict(identities, &session_id, "delete").await;
151        tracing::debug!(session_id = %session_id, "Evicted session identity on DELETE");
152    }
153}
154
155async fn evict_on_error_response(
156    identities: &McpProxyIdentityRepository,
157    response: &reqwest::Response,
158    request_headers: &HeaderMap,
159    service_name: &ServiceName,
160    method_str: &str,
161) {
162    let resp_status = response.status();
163    let header_dump: Vec<String> = response
164        .headers()
165        .iter()
166        .map(|(k, v)| format!("{}: {}", k, v.to_str().unwrap_or("?")))
167        .collect();
168    tracing::error!(
169        service = %service_name,
170        status = %resp_status,
171        headers = ?header_dump,
172        "MCP backend error response"
173    );
174
175    if resp_status == StatusCode::NOT_FOUND
176        && (method_str == "GET" || method_str == "POST")
177        && let Some(session_id) = session_id_header(request_headers)
178    {
179        evict(identities, &session_id, "stale_session").await;
180        tracing::info!(
181            service = %service_name,
182            session_id = %session_id,
183            method = %method_str,
184            reason = "stale_session",
185            "Evicted stale proxy session identity on 404; the backend no longer knows this session (restarted child?) and the client must re-initialise"
186        );
187    }
188}
189
190async fn cache_identity_from_response(
191    identities: &McpProxyIdentityRepository,
192    response: &reqwest::Response,
193    req_context: &RequestContext,
194    authenticated_user: Option<&AuthenticatedUser>,
195    service_name: &ServiceName,
196) {
197    let Some(session_id) = session_id_header(response.headers()) else {
198        return;
199    };
200    let Some(user) = authenticated_user else {
201        return;
202    };
203    let Some(auth_token) = req_context.auth_token() else {
204        return;
205    };
206    let row = ProxyIdentityRow {
207        user_id: user.id.clone(),
208        user_type: req_context.user_type(),
209        permissions: user.permissions.clone(),
210        roles: user.roles.clone(),
211        auth_token: auth_token.clone(),
212    };
213    match identities.upsert(&session_id, &row).await {
214        Ok(()) => {
215            tracing::info!(
216                service = %service_name,
217                session_id = %session_id,
218                user_id = %user.id,
219                "Stored session identity for MCP session"
220            );
221            if let Err(e) = identities
222                .attribute_session(
223                    &session_id,
224                    &McpServerId::new(service_name.as_str()),
225                    &row.user_id,
226                )
227                .await
228            {
229                tracing::warn!(
230                    service = %service_name,
231                    session_id = %session_id,
232                    error = %e,
233                    "Failed to attribute MCP session to its server and user"
234                );
235            }
236        },
237        Err(e) => tracing::error!(
238            service = %service_name,
239            session_id = %session_id,
240            user_id = %user.id,
241            error = %e,
242            "Failed to store session identity for MCP session"
243        ),
244    }
245}