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