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
13pub mod external;
14mod external_governance;
15mod external_sessions;
16mod handlers;
17pub mod mcp_session;
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::McpServerConfig;
27use systemprompt_mcp::repository::McpProxyIdentityRepository;
28use systemprompt_models::RequestContext;
29use systemprompt_runtime::AppContext;
30
31use super::auth::{AccessValidator, build_mcp_unknown_service_challenge};
32use super::backend::{HeaderInjector, ProxyError, RequestBuilder, ResponseHandler, UrlResolver};
33use super::client::ClientPool;
34use super::resolver::ServiceResolver;
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    identities: Arc<McpProxyIdentityRepository>,
53    tool_usage_repo: Option<Arc<systemprompt_mcp::repository::ToolUsageRepository>>,
54    artifact_ingest: Option<Arc<systemprompt_mcp::ArtifactIngest>>,
55}
56
57impl ProxyEngine {
58    pub fn new(identities: Arc<McpProxyIdentityRepository>) -> Self {
59        Self {
60            client_pool: ClientPool::new(),
61            identities,
62            tool_usage_repo: None,
63            artifact_ingest: None,
64        }
65    }
66
67    #[must_use]
68    pub fn with_tool_usage_repo(
69        mut self,
70        repo: Arc<systemprompt_mcp::repository::ToolUsageRepository>,
71    ) -> Self {
72        self.tool_usage_repo = Some(repo);
73        self
74    }
75
76    #[must_use]
77    pub fn with_artifact_ingest(mut self, ingest: Arc<systemprompt_mcp::ArtifactIngest>) -> Self {
78        self.artifact_ingest = Some(ingest);
79        self
80    }
81
82    pub async fn proxy_request(
83        &self,
84        target: ProxyTarget<'_>,
85        request: Request<Body>,
86        ctx: AppContext,
87    ) -> Result<Response<Body>, ProxyError> {
88        let ProxyTarget {
89            service_name,
90            path,
91            kind: proxy_kind,
92        } = target;
93        if request.extensions().get::<RequestContext>().is_none() {
94            tracing::warn!("RequestContext missing from request extensions");
95        }
96
97        if let Some(server_config) = external_server(proxy_kind, service_name, &ctx)? {
98            return self
99                .proxy_external_mcp(service_name, request, ctx, server_config)
100                .await;
101        }
102
103        let service = match ServiceResolver::resolve(service_name, &ctx).await {
104            Ok(svc) => svc,
105            Err(err) => {
106                return Err(unknown_service_error(
107                    service_name,
108                    proxy_kind,
109                    &request,
110                    &ctx,
111                    err,
112                ));
113            },
114        };
115
116        let req_ctx = request.extensions().get::<RequestContext>().cloned();
117        let authenticated_user = AccessValidator::validate(
118            request.headers(),
119            service_name,
120            &service,
121            &ctx,
122            req_ctx.as_ref(),
123        )
124        .await?;
125
126        let backend_url = UrlResolver::build_backend_url("http", "127.0.0.1", service.port, path);
127
128        let method_str = request.method().to_string();
129        let request_headers = request.headers().clone();
130        let mut headers = request_headers.clone();
131        let query = request.uri().query();
132        let full_url = UrlResolver::append_query_params(backend_url, query);
133
134        let req_context = self
135            .build_forward_context(req_ctx, &service, service_name, &request_headers)
136            .await?;
137
138        inject_forward_headers(&mut headers, &req_context, service_name);
139
140        let response = self
141            .send_to_backend(request, &full_url, &headers, service_name)
142            .await?;
143
144        if service.module_name == "mcp" {
145            mcp_session::handle_mcp_response(mcp_session::McpResponseCtx {
146                identities: &self.identities,
147                response: &response,
148                request_headers: &request_headers,
149                req_context: &req_context,
150                authenticated_user: authenticated_user.as_ref(),
151                service_name,
152                method_str: &method_str,
153            })
154            .await;
155        }
156
157        match ResponseHandler::build_response(response) {
158            Ok(resp) => Ok(resp),
159            Err(e) => {
160                tracing::error!(service = %service_name, error = %e, "Failed to build response");
161                Err(ProxyError::InvalidResponse {
162                    service: service_name.to_owned(),
163                    reason: format!("Failed to build response: {e}"),
164                })
165            },
166        }
167    }
168
169    async fn build_forward_context(
170        &self,
171        req_ctx: Option<RequestContext>,
172        service: &ServiceConfig,
173        service_name: &str,
174        request_headers: &HeaderMap,
175    ) -> Result<RequestContext, ProxyError> {
176        let mut req_context = req_ctx.ok_or_else(|| ProxyError::MissingContext {
177            message: "Request context required - proxy cannot operate without authentication"
178                .to_owned(),
179        })?;
180
181        if service.module_name == "agent" {
182            let agent_name = AgentName::try_new(service_name.to_owned()).map_err(|e| {
183                ProxyError::InvalidServiceName {
184                    service: service_name.to_owned(),
185                    reason: e.to_string(),
186                }
187            })?;
188            req_context = req_context.with_agent_name(agent_name);
189        }
190
191        if service.module_name == "mcp" && req_context.auth_token().as_str().is_empty() {
192            req_context = mcp_session::enrich_with_cached_identity(
193                &self.identities,
194                request_headers,
195                req_context,
196                service_name,
197            )
198            .await;
199        }
200
201        Ok(req_context)
202    }
203
204    async fn send_to_backend(
205        &self,
206        request: Request<Body>,
207        full_url: &str,
208        headers: &HeaderMap,
209        service_name: &str,
210    ) -> Result<reqwest::Response, ProxyError> {
211        let method_str = request.method().to_string();
212
213        let body = RequestBuilder::extract_body(request.into_body())
214            .await
215            .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
216
217        let reqwest_method = RequestBuilder::parse_method(&method_str)
218            .map_err(|reason| ProxyError::InvalidMethod { reason })?;
219
220        let client = self.client_pool.get_default_client();
221
222        let req_builder =
223            RequestBuilder::build_request(&client, reqwest_method, full_url, headers, body);
224
225        req_builder.send().await.map_err(|e| {
226            tracing::error!(service = %service_name, url = %full_url, error = %e, "Connection failed");
227            ProxyError::ConnectionFailed {
228                service: service_name.to_owned(),
229                url: full_url.to_owned(),
230                source: e,
231            }
232        })
233    }
234}
235
236fn unknown_service_error(
237    service_name: &str,
238    proxy_kind: ProxyKind,
239    request: &Request<Body>,
240    ctx: &AppContext,
241    err: ProxyError,
242) -> ProxyError {
243    if proxy_kind == ProxyKind::Mcp && matches!(err, ProxyError::ServiceNotFound { .. }) {
244        let req_ctx = request.extensions().get::<RequestContext>().cloned();
245        if let Some(challenge) = build_mcp_unknown_service_challenge(
246            service_name,
247            request.headers(),
248            ctx,
249            req_ctx.as_ref(),
250        ) {
251            return challenge;
252        }
253    }
254    err
255}
256
257fn inject_forward_headers(
258    headers: &mut HeaderMap,
259    req_context: &RequestContext,
260    service_name: &str,
261) {
262    let has_auth_before = headers.get("authorization").is_some();
263    let ctx_has_token = !req_context.auth_token().as_str().is_empty();
264
265    HeaderInjector::inject_context(headers, req_context);
266
267    let has_auth_after = headers.get("authorization").is_some();
268    tracing::debug!(
269        service = %service_name,
270        has_auth_before = has_auth_before,
271        ctx_has_token = ctx_has_token,
272        has_auth_after = has_auth_after,
273        "Proxy forwarding request"
274    );
275}
276
277fn external_server(
278    proxy_kind: ProxyKind,
279    service_name: &str,
280    ctx: &AppContext,
281) -> Result<Option<McpServerConfig>, ProxyError> {
282    if !matches!(proxy_kind, ProxyKind::Mcp) {
283        return Ok(None);
284    }
285    let found = ctx
286        .mcp_registry()
287        .find_server(service_name)
288        .map_err(|source| ProxyError::RegistryUnavailable {
289            service: service_name.to_owned(),
290            source,
291        })?;
292    Ok(found.filter(McpServerConfig::is_external))
293}