systemprompt_api/services/proxy/engine/
external.rs1use std::collections::HashMap;
13
14use axum::body::Body;
15use axum::extract::Request;
16use axum::http::{HeaderMap, HeaderName, HeaderValue, StatusCode};
17use axum::response::{IntoResponse, Response};
18use systemprompt_mcp::repository::ToolUsageRepository;
19use systemprompt_mcp::services::client::McpClient;
20use systemprompt_mcp::{McpDomainError, McpServerConfig};
21use systemprompt_models::RequestContext;
22use systemprompt_runtime::AppContext;
23
24use super::super::audit::{self, McpAudit, parse_tool_call};
25use super::super::auth::{AccessValidator, mcp_oauth_requirement};
26use super::super::backend::{ProxyError, RequestBuilder, ResponseHandler};
27use super::ProxyEngine;
28
29const MCP_PASSTHROUGH_HEADERS: [&str; 4] = [
30 "content-type",
31 "accept",
32 "mcp-session-id",
33 "mcp-protocol-version",
34];
35
36async fn authorised_context(
40 ctx: &AppContext,
41 service_name: &str,
42 headers: &HeaderMap,
43 req_ctx: Option<RequestContext>,
44) -> Result<RequestContext, ProxyError> {
45 let req_ctx = req_ctx.ok_or_else(|| ProxyError::MissingContext {
46 message: "external MCP proxy requires an authenticated request context".to_owned(),
47 })?;
48 let requirement = mcp_oauth_requirement(ctx, service_name).await?;
49 AccessValidator::validate_with_requirement(
50 headers,
51 service_name,
52 &requirement,
53 ctx,
54 Some(&req_ctx),
55 )?;
56 Ok(req_ctx)
57}
58
59impl ProxyEngine {
60 pub(super) async fn proxy_external_mcp(
61 &self,
62 service_name: &str,
63 request: Request<Body>,
64 ctx: AppContext,
65 server_config: McpServerConfig,
66 ) -> Result<Response<Body>, ProxyError> {
67 let req_ctx = authorised_context(
68 &ctx,
69 service_name,
70 request.headers(),
71 request.extensions().get::<RequestContext>().cloned(),
72 )
73 .await?;
74
75 let target = McpClient::resolve_external_proxy_target(&server_config, &req_ctx)
76 .await
77 .map_err(|e| map_resolve_error(service_name, e))?;
78
79 let method_str = request.method().to_string();
80 let incoming_headers = request.headers().clone();
81 let sessions = super::external_sessions::SessionGuard::new(
82 &self.identities,
83 service_name,
84 &req_ctx,
85 &target.headers,
86 );
87 if let Some(session) = incoming_headers.get("mcp-session-id")
88 && !sessions.accepts(session).await?
89 {
90 return Ok((
91 StatusCode::NOT_FOUND,
92 "MCP session expired; initialize again",
93 )
94 .into_response());
95 }
96 let mut body = RequestBuilder::extract_body(request.into_body())
97 .await
98 .map_err(|source| ProxyError::BodyExtractionFailed { source })?;
99 super::fixed_arguments::apply(&server_config.tools, &mut body);
100
101 super::external_governance::enforce(&ctx, &req_ctx, service_name, &body).await?;
102 let audit = build_audit(
103 self.tool_usage_repo.as_ref(),
104 self.artifact_ingest.as_ref(),
105 &req_ctx,
106 service_name,
107 &body,
108 );
109 let outbound = outbound_headers(&incoming_headers, target.headers);
110
111 let method = RequestBuilder::parse_method(&method_str)
112 .map_err(|reason| ProxyError::InvalidMethod { reason })?;
113 let client = self.client_pool.get_default_client();
114 let response = RequestBuilder::build_request(&client, method, &target.url, &outbound, body)
115 .send()
116 .await
117 .map_err(|source| ProxyError::ConnectionFailed {
118 service: service_name.to_owned(),
119 url: target.url.clone(),
120 source,
121 })?;
122
123 if method_str == "DELETE" {
124 if response.status().is_success() || response.status() == StatusCode::NOT_FOUND {
125 sessions.forget(&incoming_headers).await?;
126 }
127 } else if response.status().is_success() {
128 sessions.remember(response.headers()).await?;
129 }
130 let to_invalid = |reason| ProxyError::InvalidResponse {
131 service: service_name.to_owned(),
132 reason,
133 };
134 match audit {
135 Some(audit) => audit::record(response, audit).await.map_err(to_invalid),
136 None => ResponseHandler::build_response(response).map_err(to_invalid),
137 }
138 }
139}
140
141pub fn outbound_headers<S: std::hash::BuildHasher>(
142 incoming: &HeaderMap,
143 provider: HashMap<HeaderName, HeaderValue, S>,
144) -> HeaderMap {
145 let mut headers = HeaderMap::new();
146 for name in MCP_PASSTHROUGH_HEADERS {
147 if let Some(value) = incoming.get(name) {
148 headers.insert(HeaderName::from_static(name), value.clone());
149 }
150 }
151 for (name, value) in provider {
152 headers.insert(name, value);
153 }
154 headers
155}
156
157fn build_audit(
158 repo: Option<&std::sync::Arc<ToolUsageRepository>>,
159 ingest: Option<&std::sync::Arc<systemprompt_mcp::ArtifactIngest>>,
160 req_ctx: &RequestContext,
161 service_name: &str,
162 body: &[u8],
163) -> Option<McpAudit> {
164 let invocation = parse_tool_call(body)?;
165 let Some(repo) = repo else {
166 tracing::warn!(service = %service_name, "Tool-usage repository unavailable; external MCP call not audited");
167 return None;
168 };
169 Some(McpAudit::new(
170 std::sync::Arc::clone(repo),
171 ingest.map(std::sync::Arc::clone),
172 req_ctx.clone(),
173 service_name.to_owned(),
174 invocation,
175 ))
176}
177
178pub fn map_resolve_error(service_name: &str, error: McpDomainError) -> ProxyError {
179 match error {
180 McpDomainError::AuthRequired(_) => ProxyError::AuthenticationRequired {
181 service: service_name.to_owned(),
182 },
183 McpDomainError::ExternalAuthUnavailable { message, .. } => ProxyError::ServiceNotRunning {
184 service: service_name.to_owned(),
185 status: message,
186 },
187 other => ProxyError::InvalidResponse {
188 service: service_name.to_owned(),
189 reason: other.to_string(),
190 },
191 }
192}