Skip to main content

systemprompt_api/services/proxy/engine/
mod.rs

1//! Reverse-proxy engine for MCP and agent backends.
2//!
3//! [`ProxyEngine`] resolves a service by name, enforces access via the proxy
4//! auth boundary, forwards the request to the local backend port, and streams
5//! the response back (with SSE keep-alive). For MCP it also maintains the
6//! session-identity cache so a session-only follow-up request can be enriched
7//! with the identity established on the authenticated initialize call.
8//!
9//! Copyright (c) systemprompt.io — Business Source License 1.1.
10//! See <https://systemprompt.io> for licensing details.
11
12mod external;
13mod handlers;
14mod mcp_session;
15#[cfg(feature = "test-api")]
16pub mod test_api;
17
18use axum::body::Body;
19use axum::extract::Request;
20use axum::http::HeaderMap;
21use axum::response::Response;
22use std::collections::HashMap;
23use std::sync::Arc;
24use systemprompt_database::ServiceConfig;
25use systemprompt_identifiers::AgentName;
26use systemprompt_models::RequestContext;
27use systemprompt_runtime::AppContext;
28use tokio::sync::RwLock;
29
30use super::auth::{AccessValidator, build_mcp_unknown_service_challenge};
31use super::backend::{HeaderInjector, ProxyError, RequestBuilder, ResponseHandler, UrlResolver};
32use super::client::ClientPool;
33use super::resolver::ServiceResolver;
34use mcp_session::SessionCache;
35
36#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum ProxyKind {
38    Mcp,
39    Agent,
40}
41
42#[derive(Debug)]
43pub struct ProxyTarget<'a> {
44    pub service_name: &'a str,
45    pub path: &'a str,
46    pub kind: ProxyKind,
47}
48
49#[derive(Debug, Clone)]
50pub struct ProxyEngine {
51    client_pool: ClientPool,
52    session_cache: SessionCache,
53    tool_usage_repo: Option<Arc<systemprompt_mcp::repository::ToolUsageRepository>>,
54}
55
56impl Default for ProxyEngine {
57    fn default() -> Self {
58        Self::new()
59    }
60}
61
62impl ProxyEngine {
63    pub fn new() -> Self {
64        Self {
65            client_pool: ClientPool::new(),
66            session_cache: Arc::new(RwLock::new(HashMap::new())),
67            tool_usage_repo: None,
68        }
69    }
70
71    #[must_use]
72    pub fn with_tool_usage_repo(
73        mut self,
74        repo: Arc<systemprompt_mcp::repository::ToolUsageRepository>,
75    ) -> Self {
76        self.tool_usage_repo = Some(repo);
77        self
78    }
79
80    pub async fn proxy_request(
81        &self,
82        target: ProxyTarget<'_>,
83        request: Request<Body>,
84        ctx: AppContext,
85    ) -> Result<Response<Body>, ProxyError> {
86        let ProxyTarget {
87            service_name,
88            path,
89            kind: proxy_kind,
90        } = target;
91        if request.extensions().get::<RequestContext>().is_none() {
92            tracing::warn!("RequestContext missing from request extensions");
93        }
94
95        if matches!(proxy_kind, ProxyKind::Mcp)
96            && let Ok(Some(server_config)) = ctx.mcp_registry().find_server(service_name)
97            && server_config.is_external()
98        {
99            return self
100                .proxy_external_mcp(service_name, request, ctx, server_config)
101                .await;
102        }
103
104        let service = match ServiceResolver::resolve(service_name, &ctx).await {
105            Ok(svc) => svc,
106            Err(err) => {
107                return Err(unknown_service_error(
108                    service_name,
109                    proxy_kind,
110                    &request,
111                    &ctx,
112                    err,
113                ));
114            },
115        };
116
117        let req_ctx = request.extensions().get::<RequestContext>().cloned();
118        let authenticated_user = AccessValidator::validate(
119            request.headers(),
120            service_name,
121            &service,
122            &ctx,
123            req_ctx.as_ref(),
124        )
125        .await?;
126
127        let backend_url = UrlResolver::build_backend_url("http", "127.0.0.1", service.port, path);
128
129        let method_str = request.method().to_string();
130        let request_headers = request.headers().clone();
131        let mut headers = request_headers.clone();
132        let query = request.uri().query();
133        let full_url = UrlResolver::append_query_params(backend_url, query);
134
135        let req_context = self
136            .build_forward_context(req_ctx, &service, service_name, &request_headers)
137            .await?;
138
139        inject_forward_headers(&mut headers, &req_context, service_name);
140
141        let response = self
142            .send_to_backend(request, &full_url, &headers, service_name)
143            .await?;
144
145        if service.module_name == "mcp" {
146            mcp_session::handle_mcp_response(mcp_session::McpResponseCtx {
147                cache: &self.session_cache,
148                response: &response,
149                request_headers: &request_headers,
150                req_context: &req_context,
151                authenticated_user: authenticated_user.as_ref(),
152                service_name,
153                method_str: &method_str,
154            })
155            .await;
156        }
157
158        match ResponseHandler::build_response(response) {
159            Ok(resp) => Ok(resp),
160            Err(e) => {
161                tracing::error!(service = %service_name, error = %e, "Failed to build response");
162                Err(ProxyError::InvalidResponse {
163                    service: service_name.to_owned(),
164                    reason: format!("Failed to build response: {e}"),
165                })
166            },
167        }
168    }
169
170    async fn build_forward_context(
171        &self,
172        req_ctx: Option<RequestContext>,
173        service: &ServiceConfig,
174        service_name: &str,
175        request_headers: &HeaderMap,
176    ) -> Result<RequestContext, ProxyError> {
177        let mut req_context = req_ctx.ok_or_else(|| ProxyError::MissingContext {
178            message: "Request context required - proxy cannot operate without authentication"
179                .to_owned(),
180        })?;
181
182        // Why: an A2A service *is* the handling agent, so its name wins (the
183        // agent server enforces the same rule). An MCP server is a callee, not
184        // an agent; overwriting here would erase the caller's identity.
185        if service.module_name == "agent" {
186            let agent_name = AgentName::try_new(service_name.to_owned()).map_err(|e| {
187                ProxyError::InvalidServiceName {
188                    service: service_name.to_owned(),
189                    reason: e.to_string(),
190                }
191            })?;
192            req_context = req_context.with_agent_name(agent_name);
193        }
194
195        if service.module_name == "mcp" && req_context.auth_token().as_str().is_empty() {
196            req_context = mcp_session::enrich_with_cached_identity(
197                &self.session_cache,
198                request_headers,
199                req_context,
200                service_name,
201            )
202            .await;
203        }
204
205        Ok(req_context)
206    }
207
208    async fn send_to_backend(
209        &self,
210        request: Request<Body>,
211        full_url: &str,
212        headers: &HeaderMap,
213        service_name: &str,
214    ) -> Result<reqwest::Response, ProxyError> {
215        let method_str = request.method().to_string();
216
217        let body = RequestBuilder::extract_body(request.into_body())
218            .await
219            .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
220
221        let reqwest_method = RequestBuilder::parse_method(&method_str)
222            .map_err(|reason| ProxyError::InvalidMethod { reason })?;
223
224        let client = self.client_pool.get_default_client();
225
226        let req_builder =
227            RequestBuilder::build_request(&client, reqwest_method, full_url, headers, body);
228
229        req_builder.send().await.map_err(|e| {
230            tracing::error!(service = %service_name, url = %full_url, error = %e, "Connection failed");
231            ProxyError::ConnectionFailed {
232                service: service_name.to_owned(),
233                url: full_url.to_owned(),
234                source: e,
235            }
236        })
237    }
238}
239
240fn unknown_service_error(
241    service_name: &str,
242    proxy_kind: ProxyKind,
243    request: &Request<Body>,
244    ctx: &AppContext,
245    err: ProxyError,
246) -> ProxyError {
247    if proxy_kind == ProxyKind::Mcp && matches!(err, ProxyError::ServiceNotFound { .. }) {
248        let req_ctx = request.extensions().get::<RequestContext>().cloned();
249        if let Some(challenge) = build_mcp_unknown_service_challenge(
250            service_name,
251            request.headers(),
252            ctx,
253            req_ctx.as_ref(),
254        ) {
255            return challenge;
256        }
257    }
258    err
259}
260
261fn inject_forward_headers(
262    headers: &mut HeaderMap,
263    req_context: &RequestContext,
264    service_name: &str,
265) {
266    let has_auth_before = headers.get("authorization").is_some();
267    let ctx_has_token = !req_context.auth_token().as_str().is_empty();
268
269    HeaderInjector::inject_context(headers, req_context);
270
271    let has_auth_after = headers.get("authorization").is_some();
272    tracing::debug!(
273        service = %service_name,
274        has_auth_before = has_auth_before,
275        ctx_has_token = ctx_has_token,
276        has_auth_after = has_auth_after,
277        "Proxy forwarding request"
278    );
279}