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