systemprompt_api/services/proxy/engine/
mod.rs1mod external;
13mod handlers;
14mod mcp_session;
15#[cfg(feature = "test-api")]
16pub mod test_api;
17
18use axum::body::Body;
19use axum::extract::Request;
20use axum::http::HeaderMap;
21use axum::response::Response;
22use std::collections::HashMap;
23use std::sync::Arc;
24use systemprompt_database::ServiceConfig;
25use systemprompt_identifiers::AgentName;
26use systemprompt_models::RequestContext;
27use systemprompt_runtime::AppContext;
28use tokio::sync::RwLock;
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;
34use mcp_session::SessionCache;
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 session_cache: SessionCache,
53 tool_usage_repo: Option<Arc<systemprompt_mcp::repository::ToolUsageRepository>>,
54}
55
56impl Default for ProxyEngine {
57 fn default() -> Self {
58 Self::new()
59 }
60}
61
62impl ProxyEngine {
63 pub fn new() -> Self {
64 Self {
65 client_pool: ClientPool::new(),
66 session_cache: Arc::new(RwLock::new(HashMap::new())),
67 tool_usage_repo: None,
68 }
69 }
70
71 #[must_use]
72 pub fn with_tool_usage_repo(
73 mut self,
74 repo: Arc<systemprompt_mcp::repository::ToolUsageRepository>,
75 ) -> Self {
76 self.tool_usage_repo = Some(repo);
77 self
78 }
79
80 pub async fn proxy_request(
81 &self,
82 target: ProxyTarget<'_>,
83 request: Request<Body>,
84 ctx: AppContext,
85 ) -> Result<Response<Body>, ProxyError> {
86 let ProxyTarget {
87 service_name,
88 path,
89 kind: proxy_kind,
90 } = target;
91 if request.extensions().get::<RequestContext>().is_none() {
92 tracing::warn!("RequestContext missing from request extensions");
93 }
94
95 if matches!(proxy_kind, ProxyKind::Mcp)
96 && let Ok(Some(server_config)) = ctx.mcp_registry().find_server(service_name)
97 && server_config.is_external()
98 {
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 cache: &self.session_cache,
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" {
186 let agent_name = AgentName::try_new(service_name.to_owned()).map_err(|e| {
187 ProxyError::InvalidServiceName {
188 service: service_name.to_owned(),
189 reason: e.to_string(),
190 }
191 })?;
192 req_context = req_context.with_agent_name(agent_name);
193 }
194
195 if service.module_name == "mcp" && req_context.auth_token().as_str().is_empty() {
196 req_context = mcp_session::enrich_with_cached_identity(
197 &self.session_cache,
198 request_headers,
199 req_context,
200 service_name,
201 )
202 .await;
203 }
204
205 Ok(req_context)
206 }
207
208 async fn send_to_backend(
209 &self,
210 request: Request<Body>,
211 full_url: &str,
212 headers: &HeaderMap,
213 service_name: &str,
214 ) -> Result<reqwest::Response, ProxyError> {
215 let method_str = request.method().to_string();
216
217 let body = RequestBuilder::extract_body(request.into_body())
218 .await
219 .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
220
221 let reqwest_method = RequestBuilder::parse_method(&method_str)
222 .map_err(|reason| ProxyError::InvalidMethod { reason })?;
223
224 let client = self.client_pool.get_default_client();
225
226 let req_builder =
227 RequestBuilder::build_request(&client, reqwest_method, full_url, headers, body);
228
229 req_builder.send().await.map_err(|e| {
230 tracing::error!(service = %service_name, url = %full_url, error = %e, "Connection failed");
231 ProxyError::ConnectionFailed {
232 service: service_name.to_owned(),
233 url: full_url.to_owned(),
234 source: e,
235 }
236 })
237 }
238}
239
240fn unknown_service_error(
241 service_name: &str,
242 proxy_kind: ProxyKind,
243 request: &Request<Body>,
244 ctx: &AppContext,
245 err: ProxyError,
246) -> ProxyError {
247 if proxy_kind == ProxyKind::Mcp && matches!(err, ProxyError::ServiceNotFound { .. }) {
248 let req_ctx = request.extensions().get::<RequestContext>().cloned();
249 if let Some(challenge) = build_mcp_unknown_service_challenge(
250 service_name,
251 request.headers(),
252 ctx,
253 req_ctx.as_ref(),
254 ) {
255 return challenge;
256 }
257 }
258 err
259}
260
261fn inject_forward_headers(
262 headers: &mut HeaderMap,
263 req_context: &RequestContext,
264 service_name: &str,
265) {
266 let has_auth_before = headers.get("authorization").is_some();
267 let ctx_has_token = !req_context.auth_token().as_str().is_empty();
268
269 HeaderInjector::inject_context(headers, req_context);
270
271 let has_auth_after = headers.get("authorization").is_some();
272 tracing::debug!(
273 service = %service_name,
274 has_auth_before = has_auth_before,
275 ctx_has_token = ctx_has_token,
276 has_auth_after = has_auth_after,
277 "Proxy forwarding request"
278 );
279}