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" || service.module_name == "mcp" {
183 req_context = req_context.with_agent_name(AgentName::new(service_name.to_owned()));
184 }
185
186 if service.module_name == "mcp" && req_context.auth_token().as_str().is_empty() {
187 req_context = mcp_session::enrich_with_cached_identity(
188 &self.session_cache,
189 request_headers,
190 req_context,
191 service_name,
192 )
193 .await;
194 }
195
196 Ok(req_context)
197 }
198
199 async fn send_to_backend(
200 &self,
201 request: Request<Body>,
202 full_url: &str,
203 headers: &HeaderMap,
204 service_name: &str,
205 ) -> Result<reqwest::Response, ProxyError> {
206 let method_str = request.method().to_string();
207
208 let body = RequestBuilder::extract_body(request.into_body())
209 .await
210 .map_err(|e| ProxyError::BodyExtractionFailed { source: e })?;
211
212 let reqwest_method = RequestBuilder::parse_method(&method_str)
213 .map_err(|reason| ProxyError::InvalidMethod { reason })?;
214
215 let client = self.client_pool.get_default_client();
216
217 let req_builder =
218 RequestBuilder::build_request(&client, reqwest_method, full_url, headers, body);
219
220 req_builder.send().await.map_err(|e| {
221 tracing::error!(service = %service_name, url = %full_url, error = %e, "Connection failed");
222 ProxyError::ConnectionFailed {
223 service: service_name.to_owned(),
224 url: full_url.to_owned(),
225 source: e,
226 }
227 })
228 }
229}
230
231fn unknown_service_error(
232 service_name: &str,
233 proxy_kind: ProxyKind,
234 request: &Request<Body>,
235 ctx: &AppContext,
236 err: ProxyError,
237) -> ProxyError {
238 if proxy_kind == ProxyKind::Mcp && matches!(err, ProxyError::ServiceNotFound { .. }) {
239 let req_ctx = request.extensions().get::<RequestContext>().cloned();
240 if let Some(challenge) = build_mcp_unknown_service_challenge(
241 service_name,
242 request.headers(),
243 ctx,
244 req_ctx.as_ref(),
245 ) {
246 return challenge;
247 }
248 }
249 err
250}
251
252fn inject_forward_headers(
253 headers: &mut HeaderMap,
254 req_context: &RequestContext,
255 service_name: &str,
256) {
257 let has_auth_before = headers.get("authorization").is_some();
258 let ctx_has_token = !req_context.auth_token().as_str().is_empty();
259
260 HeaderInjector::inject_context(headers, req_context);
261
262 let has_auth_after = headers.get("authorization").is_some();
263 tracing::debug!(
264 service = %service_name,
265 has_auth_before = has_auth_before,
266 ctx_has_token = ctx_has_token,
267 has_auth_after = has_auth_after,
268 "Proxy forwarding request"
269 );
270}