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        if service.module_name == "agent" || service.module_name == "mcp" {
183            req_context = req_context.with_agent_name(AgentName::new(service_name.to_owned()));
184        }
185
186        if service.module_name == "mcp" && req_context.auth_token().as_str().is_empty() {
187            req_context = mcp_session::enrich_with_cached_identity(
188                &self.session_cache,
189                request_headers,
190                req_context,
191                service_name,
192            )
193            .await;
194        }
195
196        Ok(req_context)
197    }
198
199    async fn send_to_backend(
200        &self,
201        request: Request<Body>,
202        full_url: &str,
203        headers: &HeaderMap,
204        service_name: &str,
205    ) -> Result<reqwest::Response, ProxyError> {
206        let method_str = request.method().to_string();
207
208        let body = RequestBuilder::extract_body(request.into_body())
209            .await
210            .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
211
212        let reqwest_method = RequestBuilder::parse_method(&method_str)
213            .map_err(|reason| ProxyError::InvalidMethod { reason })?;
214
215        let client = self.client_pool.get_default_client();
216
217        let req_builder =
218            RequestBuilder::build_request(&client, reqwest_method, full_url, headers, body);
219
220        req_builder.send().await.map_err(|e| {
221            tracing::error!(service = %service_name, url = %full_url, error = %e, "Connection failed");
222            ProxyError::ConnectionFailed {
223                service: service_name.to_owned(),
224                url: full_url.to_owned(),
225                source: e,
226            }
227        })
228    }
229}
230
231fn unknown_service_error(
232    service_name: &str,
233    proxy_kind: ProxyKind,
234    request: &Request<Body>,
235    ctx: &AppContext,
236    err: ProxyError,
237) -> ProxyError {
238    if proxy_kind == ProxyKind::Mcp && matches!(err, ProxyError::ServiceNotFound { .. }) {
239        let req_ctx = request.extensions().get::<RequestContext>().cloned();
240        if let Some(challenge) = build_mcp_unknown_service_challenge(
241            service_name,
242            request.headers(),
243            ctx,
244            req_ctx.as_ref(),
245        ) {
246            return challenge;
247        }
248    }
249    err
250}
251
252fn inject_forward_headers(
253    headers: &mut HeaderMap,
254    req_context: &RequestContext,
255    service_name: &str,
256) {
257    let has_auth_before = headers.get("authorization").is_some();
258    let ctx_has_token = !req_context.auth_token().as_str().is_empty();
259
260    HeaderInjector::inject_context(headers, req_context);
261
262    let has_auth_after = headers.get("authorization").is_some();
263    tracing::debug!(
264        service = %service_name,
265        has_auth_before = has_auth_before,
266        ctx_has_token = ctx_has_token,
267        has_auth_after = has_auth_after,
268        "Proxy forwarding request"
269    );
270}