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_with_roles(
86            user_uuid,
87            String::new(),
88            String::new(),
89            identity.permissions,
90            identity.roles,
91        ))
92}
93
94#[derive(Debug)]
95pub struct McpResponseCtx<'a> {
96    pub identities: &'a McpProxyIdentityRepository,
97    pub response: &'a reqwest::Response,
98    pub request_headers: &'a HeaderMap,
99    pub req_context: &'a RequestContext,
100    pub authenticated_user: Option<&'a AuthenticatedUser>,
101    pub service_name: &'a str,
102    pub method_str: &'a str,
103}
104
105pub async fn handle_mcp_response(args: McpResponseCtx<'_>) {
106    let McpResponseCtx {
107        identities,
108        response,
109        request_headers,
110        req_context,
111        authenticated_user,
112        service_name,
113        method_str,
114    } = args;
115    let resp_status = response.status();
116    let resp_session = response
117        .headers()
118        .get("mcp-session-id")
119        .and_then(|v| v.to_str().ok())
120        .unwrap_or("none");
121    let resp_content_type = response
122        .headers()
123        .get("content-type")
124        .and_then(|v| v.to_str().ok())
125        .unwrap_or("none");
126
127    tracing::info!(
128        service = %service_name,
129        status = %resp_status,
130        resp_session_id = %resp_session,
131        content_type = %resp_content_type,
132        method = %method_str,
133        "MCP backend response"
134    );
135
136    if !resp_status.is_success() {
137        evict_on_error_response(
138            identities,
139            response,
140            request_headers,
141            service_name,
142            method_str,
143        )
144        .await;
145    }
146
147    cache_identity_from_response(
148        identities,
149        response,
150        req_context,
151        authenticated_user,
152        service_name,
153    )
154    .await;
155
156    if method_str == "DELETE"
157        && let Some(session_id) = session_id_header(request_headers)
158    {
159        evict(identities, &session_id, "delete").await;
160        tracing::debug!(session_id = %session_id, "Evicted session identity on DELETE");
161    }
162}
163
164async fn evict_on_error_response(
165    identities: &McpProxyIdentityRepository,
166    response: &reqwest::Response,
167    request_headers: &HeaderMap,
168    service_name: &str,
169    method_str: &str,
170) {
171    let resp_status = response.status();
172    let header_dump: Vec<String> = response
173        .headers()
174        .iter()
175        .map(|(k, v)| format!("{}: {}", k, v.to_str().unwrap_or("?")))
176        .collect();
177    tracing::error!(
178        service = %service_name,
179        status = %resp_status,
180        headers = ?header_dump,
181        "MCP backend error response"
182    );
183
184    if resp_status == StatusCode::NOT_FOUND
185        && (method_str == "GET" || method_str == "POST")
186        && let Some(session_id) = session_id_header(request_headers)
187    {
188        evict(identities, &session_id, "stale_session").await;
189        tracing::info!(
190            service = %service_name,
191            session_id = %session_id,
192            method = %method_str,
193            reason = "stale_session",
194            "Evicted stale proxy session identity on 404; the backend no longer knows this session (restarted child?) and the client must re-initialise"
195        );
196    }
197}
198
199async fn cache_identity_from_response(
200    identities: &McpProxyIdentityRepository,
201    response: &reqwest::Response,
202    req_context: &RequestContext,
203    authenticated_user: Option<&AuthenticatedUser>,
204    service_name: &str,
205) {
206    let Some(session_id) = session_id_header(response.headers()) else {
207        return;
208    };
209    let Some(user) = authenticated_user else {
210        return;
211    };
212    let row = ProxyIdentityRow {
213        user_id: UserId::new(user.id.to_string()),
214        user_type: req_context.user_type(),
215        permissions: user.permissions.clone(),
216        roles: user.roles.clone(),
217        auth_token: req_context.auth_token().clone(),
218    };
219    match identities.upsert(&session_id, &row).await {
220        Ok(()) => tracing::info!(
221            service = %service_name,
222            session_id = %session_id,
223            user_id = %user.id,
224            "Stored session identity for MCP session"
225        ),
226        Err(e) => tracing::error!(
227            service = %service_name,
228            session_id = %session_id,
229            user_id = %user.id,
230            error = %e,
231            "Failed to store session identity for MCP session"
232        ),
233    }
234}