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