systemprompt_api/services/proxy/engine/
mod.rs1pub 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;
26use systemprompt_identifiers::AgentName;
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 str,
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 tool_usage_repo: Option<Arc<systemprompt_mcp::repository::ToolUsageRepository>>,
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 tool_usage_repo: 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 ) -> Self {
73 self.tool_usage_repo = Some(repo);
74 self
75 }
76
77 #[must_use]
78 pub fn with_artifact_ingest(mut self, ingest: Arc<systemprompt_mcp::ArtifactIngest>) -> Self {
79 self.artifact_ingest = Some(ingest);
80 self
81 }
82
83 pub async fn proxy_request(
84 &self,
85 target: ProxyTarget<'_>,
86 request: Request<Body>,
87 ctx: AppContext,
88 ) -> Result<Response<Body>, ProxyError> {
89 let ProxyTarget {
90 service_name,
91 path,
92 kind: proxy_kind,
93 } = target;
94 if request.extensions().get::<RequestContext>().is_none() {
95 tracing::warn!("RequestContext missing from request extensions");
96 }
97
98 if let Some(server_config) = external_server(proxy_kind, service_name, &ctx)? {
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 identities: &self.identities,
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" {
183 let agent_name = AgentName::try_new(service_name.to_owned()).map_err(|e| {
184 ProxyError::InvalidServiceName {
185 service: service_name.to_owned(),
186 reason: e.to_string(),
187 }
188 })?;
189 req_context = req_context.with_agent_name(agent_name);
190 }
191
192 if service.module_name == "mcp" && req_context.auth_token().as_str().is_empty() {
193 req_context = mcp_session::enrich_with_cached_identity(
194 &self.identities,
195 request_headers,
196 req_context,
197 service_name,
198 )
199 .await;
200 }
201
202 Ok(req_context)
203 }
204
205 async fn send_to_backend(
206 &self,
207 request: Request<Body>,
208 full_url: &str,
209 headers: &HeaderMap,
210 service_name: &str,
211 ) -> Result<reqwest::Response, ProxyError> {
212 let method_str = request.method().to_string();
213
214 let body = RequestBuilder::extract_body(request.into_body())
215 .await
216 .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
217
218 let reqwest_method = RequestBuilder::parse_method(&method_str)
219 .map_err(|reason| ProxyError::InvalidMethod { reason })?;
220
221 let client = self.client_pool.get_default_client();
222
223 let req_builder =
224 RequestBuilder::build_request(&client, reqwest_method, full_url, headers, body);
225
226 req_builder.send().await.map_err(|e| {
227 tracing::error!(service = %service_name, url = %full_url, error = %e, "Connection failed");
228 ProxyError::ConnectionFailed {
229 service: service_name.to_owned(),
230 url: full_url.to_owned(),
231 source: e,
232 }
233 })
234 }
235}
236
237fn unknown_service_error(
238 service_name: &str,
239 proxy_kind: ProxyKind,
240 request: &Request<Body>,
241 ctx: &AppContext,
242 err: ProxyError,
243) -> ProxyError {
244 if proxy_kind == ProxyKind::Mcp && matches!(err, ProxyError::ServiceNotFound { .. }) {
245 let req_ctx = request.extensions().get::<RequestContext>().cloned();
246 if let Some(challenge) = build_mcp_unknown_service_challenge(
247 service_name,
248 request.headers(),
249 ctx,
250 req_ctx.as_ref(),
251 ) {
252 return challenge;
253 }
254 }
255 err
256}
257
258fn inject_forward_headers(
259 headers: &mut HeaderMap,
260 req_context: &RequestContext,
261 service_name: &str,
262) {
263 let has_auth_before = headers.get("authorization").is_some();
264 let ctx_has_token = !req_context.auth_token().as_str().is_empty();
265
266 HeaderInjector::inject_context(headers, req_context);
267
268 let has_auth_after = headers.get("authorization").is_some();
269 tracing::debug!(
270 service = %service_name,
271 has_auth_before = has_auth_before,
272 ctx_has_token = ctx_has_token,
273 has_auth_after = has_auth_after,
274 "Proxy forwarding request"
275 );
276}
277
278fn external_server(
279 proxy_kind: ProxyKind,
280 service_name: &str,
281 ctx: &AppContext,
282) -> Result<Option<McpServerConfig>, ProxyError> {
283 if !matches!(proxy_kind, ProxyKind::Mcp) {
284 return Ok(None);
285 }
286 let found = ctx
287 .mcp_registry()
288 .find_server(service_name)
289 .map_err(|source| ProxyError::RegistryUnavailable {
290 service: service_name.to_owned(),
291 source,
292 })?;
293 Ok(found.filter(McpServerConfig::is_external))
294}