systemprompt_api/services/proxy/engine/
mod.rs1pub mod external;
14mod external_governance;
15mod external_sessions;
16mod handlers;
17pub mod mcp_session;
18
19use axum::body::Body;
20use axum::extract::Request;
21use axum::http::HeaderMap;
22use axum::response::Response;
23use std::sync::Arc;
24use systemprompt_database::ServiceConfig;
25use systemprompt_identifiers::AgentName;
26use systemprompt_mcp::McpServerConfig;
27use systemprompt_mcp::repository::McpProxyIdentityRepository;
28use systemprompt_models::RequestContext;
29use systemprompt_runtime::AppContext;
30
31use super::auth::{AccessValidator, build_mcp_unknown_service_challenge};
32use super::backend::{HeaderInjector, ProxyError, RequestBuilder, ResponseHandler, UrlResolver};
33use super::client::ClientPool;
34use super::resolver::ServiceResolver;
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 identities: Arc<McpProxyIdentityRepository>,
53 tool_usage_repo: Option<Arc<systemprompt_mcp::repository::ToolUsageRepository>>,
54 artifact_ingest: Option<Arc<systemprompt_mcp::ArtifactIngest>>,
55}
56
57impl ProxyEngine {
58 pub fn new(identities: Arc<McpProxyIdentityRepository>) -> Self {
59 Self {
60 client_pool: ClientPool::new(),
61 identities,
62 tool_usage_repo: None,
63 artifact_ingest: None,
64 }
65 }
66
67 #[must_use]
68 pub fn with_tool_usage_repo(
69 mut self,
70 repo: Arc<systemprompt_mcp::repository::ToolUsageRepository>,
71 ) -> Self {
72 self.tool_usage_repo = Some(repo);
73 self
74 }
75
76 #[must_use]
77 pub fn with_artifact_ingest(mut self, ingest: Arc<systemprompt_mcp::ArtifactIngest>) -> Self {
78 self.artifact_ingest = Some(ingest);
79 self
80 }
81
82 pub async fn proxy_request(
83 &self,
84 target: ProxyTarget<'_>,
85 request: Request<Body>,
86 ctx: AppContext,
87 ) -> Result<Response<Body>, ProxyError> {
88 let ProxyTarget {
89 service_name,
90 path,
91 kind: proxy_kind,
92 } = target;
93 if request.extensions().get::<RequestContext>().is_none() {
94 tracing::warn!("RequestContext missing from request extensions");
95 }
96
97 if let Some(server_config) = external_server(proxy_kind, service_name, &ctx)? {
98 return self
99 .proxy_external_mcp(service_name, request, ctx, server_config)
100 .await;
101 }
102
103 let service = match ServiceResolver::resolve(service_name, &ctx).await {
104 Ok(svc) => svc,
105 Err(err) => {
106 return Err(unknown_service_error(
107 service_name,
108 proxy_kind,
109 &request,
110 &ctx,
111 err,
112 ));
113 },
114 };
115
116 let req_ctx = request.extensions().get::<RequestContext>().cloned();
117 let authenticated_user = AccessValidator::validate(
118 request.headers(),
119 service_name,
120 &service,
121 &ctx,
122 req_ctx.as_ref(),
123 )
124 .await?;
125
126 let backend_url = UrlResolver::build_backend_url("http", "127.0.0.1", service.port, path);
127
128 let method_str = request.method().to_string();
129 let request_headers = request.headers().clone();
130 let mut headers = request_headers.clone();
131 let query = request.uri().query();
132 let full_url = UrlResolver::append_query_params(backend_url, query);
133
134 let req_context = self
135 .build_forward_context(req_ctx, &service, service_name, &request_headers)
136 .await?;
137
138 inject_forward_headers(&mut headers, &req_context, service_name);
139
140 let response = self
141 .send_to_backend(request, &full_url, &headers, service_name)
142 .await?;
143
144 if service.module_name == "mcp" {
145 mcp_session::handle_mcp_response(mcp_session::McpResponseCtx {
146 identities: &self.identities,
147 response: &response,
148 request_headers: &request_headers,
149 req_context: &req_context,
150 authenticated_user: authenticated_user.as_ref(),
151 service_name,
152 method_str: &method_str,
153 })
154 .await;
155 }
156
157 match ResponseHandler::build_response(response) {
158 Ok(resp) => Ok(resp),
159 Err(e) => {
160 tracing::error!(service = %service_name, error = %e, "Failed to build response");
161 Err(ProxyError::InvalidResponse {
162 service: service_name.to_owned(),
163 reason: format!("Failed to build response: {e}"),
164 })
165 },
166 }
167 }
168
169 async fn build_forward_context(
170 &self,
171 req_ctx: Option<RequestContext>,
172 service: &ServiceConfig,
173 service_name: &str,
174 request_headers: &HeaderMap,
175 ) -> Result<RequestContext, ProxyError> {
176 let mut req_context = req_ctx.ok_or_else(|| ProxyError::MissingContext {
177 message: "Request context required - proxy cannot operate without authentication"
178 .to_owned(),
179 })?;
180
181 if service.module_name == "agent" {
182 let agent_name = AgentName::try_new(service_name.to_owned()).map_err(|e| {
183 ProxyError::InvalidServiceName {
184 service: service_name.to_owned(),
185 reason: e.to_string(),
186 }
187 })?;
188 req_context = req_context.with_agent_name(agent_name);
189 }
190
191 if service.module_name == "mcp" && req_context.auth_token().as_str().is_empty() {
192 req_context = mcp_session::enrich_with_cached_identity(
193 &self.identities,
194 request_headers,
195 req_context,
196 service_name,
197 )
198 .await;
199 }
200
201 Ok(req_context)
202 }
203
204 async fn send_to_backend(
205 &self,
206 request: Request<Body>,
207 full_url: &str,
208 headers: &HeaderMap,
209 service_name: &str,
210 ) -> Result<reqwest::Response, ProxyError> {
211 let method_str = request.method().to_string();
212
213 let body = RequestBuilder::extract_body(request.into_body())
214 .await
215 .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
216
217 let reqwest_method = RequestBuilder::parse_method(&method_str)
218 .map_err(|reason| ProxyError::InvalidMethod { reason })?;
219
220 let client = self.client_pool.get_default_client();
221
222 let req_builder =
223 RequestBuilder::build_request(&client, reqwest_method, full_url, headers, body);
224
225 req_builder.send().await.map_err(|e| {
226 tracing::error!(service = %service_name, url = %full_url, error = %e, "Connection failed");
227 ProxyError::ConnectionFailed {
228 service: service_name.to_owned(),
229 url: full_url.to_owned(),
230 source: e,
231 }
232 })
233 }
234}
235
236fn unknown_service_error(
237 service_name: &str,
238 proxy_kind: ProxyKind,
239 request: &Request<Body>,
240 ctx: &AppContext,
241 err: ProxyError,
242) -> ProxyError {
243 if proxy_kind == ProxyKind::Mcp && matches!(err, ProxyError::ServiceNotFound { .. }) {
244 let req_ctx = request.extensions().get::<RequestContext>().cloned();
245 if let Some(challenge) = build_mcp_unknown_service_challenge(
246 service_name,
247 request.headers(),
248 ctx,
249 req_ctx.as_ref(),
250 ) {
251 return challenge;
252 }
253 }
254 err
255}
256
257fn inject_forward_headers(
258 headers: &mut HeaderMap,
259 req_context: &RequestContext,
260 service_name: &str,
261) {
262 let has_auth_before = headers.get("authorization").is_some();
263 let ctx_has_token = !req_context.auth_token().as_str().is_empty();
264
265 HeaderInjector::inject_context(headers, req_context);
266
267 let has_auth_after = headers.get("authorization").is_some();
268 tracing::debug!(
269 service = %service_name,
270 has_auth_before = has_auth_before,
271 ctx_has_token = ctx_has_token,
272 has_auth_after = has_auth_after,
273 "Proxy forwarding request"
274 );
275}
276
277fn external_server(
278 proxy_kind: ProxyKind,
279 service_name: &str,
280 ctx: &AppContext,
281) -> Result<Option<McpServerConfig>, ProxyError> {
282 if !matches!(proxy_kind, ProxyKind::Mcp) {
283 return Ok(None);
284 }
285 let found = ctx
286 .mcp_registry()
287 .find_server(service_name)
288 .map_err(|source| ProxyError::RegistryUnavailable {
289 service: service_name.to_owned(),
290 source,
291 })?;
292 Ok(found.filter(McpServerConfig::is_external))
293}