systemprompt_api/services/middleware/context/middleware/
flavours.rs1use std::sync::Arc;
11
12use axum::extract::Request;
13use axum::middleware::Next;
14use axum::response::Response;
15use systemprompt_identifiers::{AgentName, ContextId};
16use systemprompt_models::execution::context::RequestContext;
17use systemprompt_security::HeaderExtractor;
18use tracing::Instrument;
19
20use super::super::extractors::ContextExtractor;
21use super::error::log_error_response;
22use super::support::{DynExtractor, create_request_span, session_context_required_error};
23
24#[derive(Clone, Copy, Debug, Default)]
31pub struct PublicContextMiddleware;
32
33impl PublicContextMiddleware {
34 #[must_use]
35 pub const fn new() -> Self {
36 Self
37 }
38
39 pub async fn handle(&self, mut request: Request, next: Next) -> Response {
40 let Some(mut req_ctx) = request.extensions().get::<RequestContext>().cloned() else {
41 let trace_id = HeaderExtractor::extract_trace_id(request.headers());
42 let path = request.uri().path().to_owned();
43 let method = request.method().to_string();
44 return session_context_required_error(&trace_id, &path, &method);
45 };
46
47 let headers = request.headers();
48 if let Some(context_id) = headers.get("x-context-id")
49 && let Ok(id) = context_id.to_str()
50 {
51 req_ctx.execution.context_id = ContextId::new(id.to_owned());
52 }
53
54 if let Some(agent_name) = headers.get("x-agent-name")
55 && let Ok(name) = agent_name.to_str()
56 {
57 req_ctx.execution.agent_name = AgentName::new(name.to_owned());
58 }
59
60 let span = create_request_span(&req_ctx);
61 request.extensions_mut().insert(req_ctx);
62 next.run(request).instrument(span).await
63 }
64}
65
66#[derive(Clone)]
70pub struct UserOnlyContextMiddleware {
71 extractor: DynExtractor,
72}
73
74impl std::fmt::Debug for UserOnlyContextMiddleware {
75 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
76 f.debug_struct("UserOnlyContextMiddleware").finish()
77 }
78}
79
80impl UserOnlyContextMiddleware {
81 pub fn new<E>(extractor: E) -> Self
82 where
83 E: ContextExtractor + Send + Sync + 'static,
84 {
85 Self {
86 extractor: Arc::new(extractor),
87 }
88 }
89
90 pub async fn handle(&self, mut request: Request, next: Next) -> Response {
91 let trace_id = HeaderExtractor::extract_trace_id(request.headers());
92 let path = request.uri().path().to_owned();
93 let method = request.method().to_string();
94
95 match self.extractor.extract_from_headers(request.headers()).await {
96 Ok(context) => {
97 let span = create_request_span(&context);
98 request.extensions_mut().insert(context);
99 next.run(request).instrument(span).await
100 },
101 Err(e) => log_error_response(&e, &trace_id, &path, &method),
102 }
103 }
104}
105
106#[derive(Clone)]
112pub struct A2AContextMiddleware {
113 extractor: DynExtractor,
114}
115
116impl std::fmt::Debug for A2AContextMiddleware {
117 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
118 f.debug_struct("A2AContextMiddleware").finish()
119 }
120}
121
122impl A2AContextMiddleware {
123 pub fn new<E>(extractor: E) -> Self
124 where
125 E: ContextExtractor + Send + Sync + 'static,
126 {
127 Self {
128 extractor: Arc::new(extractor),
129 }
130 }
131
132 pub async fn handle(&self, request: Request, next: Next) -> Response {
133 let trace_id = HeaderExtractor::extract_trace_id(request.headers());
134 let path = request.uri().path().to_owned();
135 let method = request.method().to_string();
136
137 match self.extractor.extract_from_request(request).await {
138 Ok((context, reconstructed_request)) => {
139 let span = create_request_span(&context);
140 let mut req = reconstructed_request;
141 req.extensions_mut().insert(context);
142 next.run(req).instrument(span).await
143 },
144 Err(e) => log_error_response(&e, &trace_id, &path, &method),
145 }
146 }
147}
148
149#[derive(Clone)]
161pub struct McpContextMiddleware {
162 extractor: DynExtractor,
163}
164
165impl std::fmt::Debug for McpContextMiddleware {
166 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
167 f.debug_struct("McpContextMiddleware").finish()
168 }
169}
170
171impl McpContextMiddleware {
172 pub fn new<E>(extractor: E) -> Self
173 where
174 E: ContextExtractor + Send + Sync + 'static,
175 {
176 Self {
177 extractor: Arc::new(extractor),
178 }
179 }
180
181 pub async fn handle(&self, request: Request, next: Next) -> Response {
182 let trace_id = HeaderExtractor::extract_trace_id(request.headers());
183 let path = request.uri().path().to_owned();
184 let method = request.method().to_string();
185
186 match self.extractor.extract_from_headers(request.headers()).await {
187 Ok(context) => {
188 let span = create_request_span(&context);
189 let mut req = request;
190 req.extensions_mut().insert(context);
191 next.run(req).instrument(span).await
192 },
193 Err(e) => {
194 if let Some(ctx) = request.extensions().get::<RequestContext>().cloned() {
195 tracing::debug!(
196 error = %e,
197 trace_id = %trace_id,
198 "MCP header extraction failed, using session context"
199 );
200 let span = create_request_span(&ctx);
201 next.run(request).instrument(span).await
202 } else {
203 session_context_required_error(&trace_id, &path, &method)
204 }
205 },
206 }
207 }
208}