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, ServiceModule};
26use systemprompt_identifiers::{AgentName, ServiceName};
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 ServiceName,
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    intent_claims: Option<systemprompt_mcp::IntentClaimService>,
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            intent_claims: 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        intents: systemprompt_traits::DynToolCallIntentClaims,
73    ) -> Self {
74        self.intent_claims = Some(systemprompt_mcp::IntentClaimService::new(intents, repo));
75        self
76    }
77
78    #[must_use]
79    pub fn with_artifact_ingest(mut self, ingest: Arc<systemprompt_mcp::ArtifactIngest>) -> Self {
80        self.artifact_ingest = Some(ingest);
81        self
82    }
83
84    pub async fn proxy_request(
85        &self,
86        target: ProxyTarget<'_>,
87        request: Request<Body>,
88        ctx: AppContext,
89    ) -> Result<Response<Body>, ProxyError> {
90        let ProxyTarget {
91            service_name,
92            path,
93            kind: proxy_kind,
94        } = target;
95        if request.extensions().get::<RequestContext>().is_none() {
96            tracing::warn!("RequestContext missing from request extensions");
97        }
98
99        if let Some(server_config) = external_server(proxy_kind, service_name, &ctx)? {
100            return self
101                .proxy_external_mcp(service_name, request, ctx, server_config)
102                .await;
103        }
104
105        let service = match ServiceResolver::resolve(service_name, &ctx).await {
106            Ok(svc) => svc,
107            Err(err) => {
108                return Err(unknown_service_error(
109                    service_name,
110                    proxy_kind,
111                    &request,
112                    &ctx,
113                    err,
114                ));
115            },
116        };
117
118        let req_ctx = request.extensions().get::<RequestContext>().cloned();
119        let authenticated_user = AccessValidator::validate(
120            request.headers(),
121            service_name,
122            &service,
123            &ctx,
124            req_ctx.as_ref(),
125        )
126        .await?;
127
128        let backend_url = UrlResolver::build_backend_url("http", "127.0.0.1", service.port, path);
129
130        let method_str = request.method().to_string();
131        let request_headers = request.headers().clone();
132        let mut headers = request_headers.clone();
133        let query = request.uri().query();
134        let full_url = UrlResolver::append_query_params(backend_url, query);
135
136        let req_context = self
137            .build_forward_context(req_ctx, &service, service_name, &request_headers)
138            .await?;
139
140        inject_forward_headers(&mut headers, &req_context, service_name);
141
142        let response = self
143            .send_to_backend(request, &full_url, &headers, service_name)
144            .await?;
145
146        if service.module_name == ServiceModule::Mcp {
147            mcp_session::handle_mcp_response(mcp_session::McpResponseCtx {
148                identities: &self.identities,
149                response: &response,
150                request_headers: &request_headers,
151                req_context: &req_context,
152                authenticated_user: authenticated_user.as_ref(),
153                service_name,
154                method_str: &method_str,
155            })
156            .await;
157        }
158
159        ResponseHandler::build_response(response).map_err(|source| ProxyError::InvalidResponse {
160            service: service_name.to_string(),
161            source,
162        })
163    }
164
165    async fn build_forward_context(
166        &self,
167        req_ctx: Option<RequestContext>,
168        service: &ServiceConfig,
169        service_name: &ServiceName,
170        request_headers: &HeaderMap,
171    ) -> Result<RequestContext, ProxyError> {
172        let mut req_context = req_ctx.ok_or_else(|| ProxyError::MissingContext {
173            message: "Request context required - proxy cannot operate without authentication"
174                .to_owned(),
175        })?;
176
177        if service.module_name == ServiceModule::Agent {
178            let agent_name = AgentName::try_new(service_name.as_str()).map_err(|source| {
179                ProxyError::InvalidServiceName {
180                    service: service_name.to_string(),
181                    source,
182                }
183            })?;
184            req_context = req_context.with_agent_name(agent_name);
185        }
186
187        if service.module_name == ServiceModule::Mcp && req_context.auth_token().is_none() {
188            req_context = mcp_session::enrich_with_cached_identity(
189                &self.identities,
190                request_headers,
191                req_context,
192                service_name,
193            )
194            .await;
195        }
196
197        Ok(req_context)
198    }
199
200    async fn send_to_backend(
201        &self,
202        request: Request<Body>,
203        full_url: &str,
204        headers: &HeaderMap,
205        service_name: &ServiceName,
206    ) -> Result<reqwest::Response, ProxyError> {
207        let method_str = request.method().to_string();
208
209        let body = RequestBuilder::extract_body(request.into_body())
210            .await
211            .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
212
213        let reqwest_method = RequestBuilder::parse_method(&method_str)?;
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
221            .send()
222            .await
223            .map_err(|e| ProxyError::ConnectionFailed {
224                service: service_name.to_string(),
225                url: full_url.to_owned(),
226                source: e,
227            })
228    }
229}
230
231fn unknown_service_error(
232    service_name: &ServiceName,
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: &ServiceName,
256) {
257    let has_auth_before = headers.get("authorization").is_some();
258    let ctx_has_token = req_context.auth_token().is_some();
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}
271
272fn external_server(
273    proxy_kind: ProxyKind,
274    service_name: &ServiceName,
275    ctx: &AppContext,
276) -> Result<Option<McpServerConfig>, ProxyError> {
277    if !matches!(proxy_kind, ProxyKind::Mcp) {
278        return Ok(None);
279    }
280    let found = ctx
281        .mcp_registry()
282        .find_server(service_name.as_str())
283        .map_err(|source| ProxyError::RegistryUnavailable {
284            service: service_name.to_string(),
285            source,
286        })?;
287    Ok(found.filter(McpServerConfig::is_external))
288}