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;
16pub mod fixed_arguments;
17mod handlers;
18pub mod mcp_session;
19
20use axum::body::Body;
21use axum::extract::Request;
22use axum::http::HeaderMap;
23use axum::response::Response;
24use std::sync::Arc;
25use systemprompt_database::ServiceConfig;
26use systemprompt_identifiers::AgentName;
27use systemprompt_mcp::McpServerConfig;
28use systemprompt_mcp::repository::McpProxyIdentityRepository;
29use systemprompt_models::RequestContext;
30use systemprompt_runtime::AppContext;
31
32use super::auth::{AccessValidator, build_mcp_unknown_service_challenge};
33use super::backend::{HeaderInjector, ProxyError, RequestBuilder, ResponseHandler, UrlResolver};
34use super::client::ClientPool;
35use super::resolver::ServiceResolver;
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq)]
38pub enum ProxyKind {
39    Mcp,
40    Agent,
41}
42
43#[derive(Debug)]
44pub struct ProxyTarget<'a> {
45    pub service_name: &'a str,
46    pub path: &'a str,
47    pub kind: ProxyKind,
48}
49
50#[derive(Debug, Clone)]
51pub struct ProxyEngine {
52    client_pool: ClientPool,
53    identities: Arc<McpProxyIdentityRepository>,
54    tool_usage_repo: Option<Arc<systemprompt_mcp::repository::ToolUsageRepository>>,
55    artifact_ingest: Option<Arc<systemprompt_mcp::ArtifactIngest>>,
56}
57
58impl ProxyEngine {
59    pub fn new(identities: Arc<McpProxyIdentityRepository>) -> Self {
60        Self {
61            client_pool: ClientPool::new(),
62            identities,
63            tool_usage_repo: None,
64            artifact_ingest: None,
65        }
66    }
67
68    #[must_use]
69    pub fn with_tool_usage_repo(
70        mut self,
71        repo: Arc<systemprompt_mcp::repository::ToolUsageRepository>,
72    ) -> Self {
73        self.tool_usage_repo = Some(repo);
74        self
75    }
76
77    #[must_use]
78    pub fn with_artifact_ingest(mut self, ingest: Arc<systemprompt_mcp::ArtifactIngest>) -> Self {
79        self.artifact_ingest = Some(ingest);
80        self
81    }
82
83    pub async fn proxy_request(
84        &self,
85        target: ProxyTarget<'_>,
86        request: Request<Body>,
87        ctx: AppContext,
88    ) -> Result<Response<Body>, ProxyError> {
89        let ProxyTarget {
90            service_name,
91            path,
92            kind: proxy_kind,
93        } = target;
94        if request.extensions().get::<RequestContext>().is_none() {
95            tracing::warn!("RequestContext missing from request extensions");
96        }
97
98        if let Some(server_config) = external_server(proxy_kind, service_name, &ctx)? {
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                identities: &self.identities,
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        if service.module_name == "agent" {
183            let agent_name = AgentName::try_new(service_name.to_owned()).map_err(|e| {
184                ProxyError::InvalidServiceName {
185                    service: service_name.to_owned(),
186                    reason: e.to_string(),
187                }
188            })?;
189            req_context = req_context.with_agent_name(agent_name);
190        }
191
192        if service.module_name == "mcp" && req_context.auth_token().as_str().is_empty() {
193            req_context = mcp_session::enrich_with_cached_identity(
194                &self.identities,
195                request_headers,
196                req_context,
197                service_name,
198            )
199            .await;
200        }
201
202        Ok(req_context)
203    }
204
205    async fn send_to_backend(
206        &self,
207        request: Request<Body>,
208        full_url: &str,
209        headers: &HeaderMap,
210        service_name: &str,
211    ) -> Result<reqwest::Response, ProxyError> {
212        let method_str = request.method().to_string();
213
214        let body = RequestBuilder::extract_body(request.into_body())
215            .await
216            .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
217
218        let reqwest_method = RequestBuilder::parse_method(&method_str)
219            .map_err(|reason| ProxyError::InvalidMethod { reason })?;
220
221        let client = self.client_pool.get_default_client();
222
223        let req_builder =
224            RequestBuilder::build_request(&client, reqwest_method, full_url, headers, body);
225
226        req_builder.send().await.map_err(|e| {
227            tracing::error!(service = %service_name, url = %full_url, error = %e, "Connection failed");
228            ProxyError::ConnectionFailed {
229                service: service_name.to_owned(),
230                url: full_url.to_owned(),
231                source: e,
232            }
233        })
234    }
235}
236
237fn unknown_service_error(
238    service_name: &str,
239    proxy_kind: ProxyKind,
240    request: &Request<Body>,
241    ctx: &AppContext,
242    err: ProxyError,
243) -> ProxyError {
244    if proxy_kind == ProxyKind::Mcp && matches!(err, ProxyError::ServiceNotFound { .. }) {
245        let req_ctx = request.extensions().get::<RequestContext>().cloned();
246        if let Some(challenge) = build_mcp_unknown_service_challenge(
247            service_name,
248            request.headers(),
249            ctx,
250            req_ctx.as_ref(),
251        ) {
252            return challenge;
253        }
254    }
255    err
256}
257
258fn inject_forward_headers(
259    headers: &mut HeaderMap,
260    req_context: &RequestContext,
261    service_name: &str,
262) {
263    let has_auth_before = headers.get("authorization").is_some();
264    let ctx_has_token = !req_context.auth_token().as_str().is_empty();
265
266    HeaderInjector::inject_context(headers, req_context);
267
268    let has_auth_after = headers.get("authorization").is_some();
269    tracing::debug!(
270        service = %service_name,
271        has_auth_before = has_auth_before,
272        ctx_has_token = ctx_has_token,
273        has_auth_after = has_auth_after,
274        "Proxy forwarding request"
275    );
276}
277
278fn external_server(
279    proxy_kind: ProxyKind,
280    service_name: &str,
281    ctx: &AppContext,
282) -> Result<Option<McpServerConfig>, ProxyError> {
283    if !matches!(proxy_kind, ProxyKind::Mcp) {
284        return Ok(None);
285    }
286    let found = ctx
287        .mcp_registry()
288        .find_server(service_name)
289        .map_err(|source| ProxyError::RegistryUnavailable {
290            service: service_name.to_owned(),
291            source,
292        })?;
293    Ok(found.filter(McpServerConfig::is_external))
294}