systemprompt_api/services/middleware/jwt/
context.rs1use async_trait::async_trait;
13use axum::body::Body;
14use axum::extract::Request;
15use axum::http::HeaderMap;
16use std::sync::Arc;
17
18use crate::services::middleware::context::ContextExtractor;
19use systemprompt_identifiers::ContextId;
20use systemprompt_models::execution::context::{ContextExtractionError, RequestContext};
21use systemprompt_security::{JwtUserContext, TokenExtractor, extract_user_context};
22use systemprompt_traits::{SessionProvider, UserProvider};
23
24use super::params::{BuildContextParams, build_context, extract_common_headers};
25use super::revocation::JtiRevocationChecker;
26use super::validation::{UserCache, user_is_admin, validate_session_exists, validate_user_exists};
27
28#[derive(Clone)]
29pub struct JwtContextExtractor {
30 token_extractor: TokenExtractor,
31 session_provider: Arc<dyn SessionProvider>,
32 user_provider: Arc<dyn UserProvider>,
33 user_cache: Arc<UserCache>,
34 jti_revocation: JtiRevocationChecker,
35 issuer: String,
36}
37
38impl std::fmt::Debug for JwtContextExtractor {
39 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
40 f.debug_struct("JwtContextExtractor")
41 .field("token_extractor", &self.token_extractor)
42 .finish_non_exhaustive()
43 }
44}
45
46impl JwtContextExtractor {
47 pub fn new(
48 session_provider: Arc<dyn SessionProvider>,
49 user_provider: Arc<dyn UserProvider>,
50 jti_revocation: JtiRevocationChecker,
51 issuer: String,
52 ) -> Self {
53 Self {
54 token_extractor: TokenExtractor::browser_only(),
55 session_provider,
56 user_provider,
57 user_cache: UserCache::new(),
58 jti_revocation,
59 issuer,
60 }
61 }
62
63 fn extract_jwt_context(
64 &self,
65 headers: &HeaderMap,
66 ) -> Result<JwtUserContext, ContextExtractionError> {
67 let token = self
68 .token_extractor
69 .extract(headers)
70 .map_err(|_e| ContextExtractionError::MissingAuthHeader)?;
71 extract_user_context(&token, &self.issuer)
72 .map_err(|e| ContextExtractionError::InvalidToken(e.into()))
73 }
74
75 async fn validate(
76 &self,
77 jwt_context: &JwtUserContext,
78 route_context: &str,
79 ) -> Result<systemprompt_traits::AuthUser, ContextExtractionError> {
80 if jwt_context.session_id.as_str().is_empty() {
81 return Err(ContextExtractionError::MissingSessionId);
82 }
83 if jwt_context.user_id.as_str().is_empty() {
84 return Err(ContextExtractionError::MissingUserId);
85 }
86 let validated = validate_user_exists(
87 &self.user_provider,
88 &self.user_cache,
89 jwt_context,
90 route_context,
91 )
92 .await?;
93 validate_session_exists(&self.session_provider, jwt_context, route_context).await?;
94 if let Some(jti) = &jwt_context.jti {
95 self.jti_revocation.ensure_not_revoked(jti).await?;
96 }
97 Ok(validated.user)
98 }
99
100 pub async fn extract_standard(
101 &self,
102 headers: &HeaderMap,
103 ) -> Result<RequestContext, ContextExtractionError> {
104 let jwt_context = self.extract_jwt_context(headers)?;
105 let user = self.validate(&jwt_context, "").await?;
106
107 let context_id = headers
108 .get("x-context-id")
109 .and_then(|h| h.to_str().ok())
110 .filter(|s| !s.is_empty())
111 .and_then(|s| ContextId::try_new(s).ok())
112 .unwrap_or_else(|| ContextId::derived_from_session(&jwt_context.session_id));
113
114 let (trace_id, task_id, auth_token, agent_name) =
115 extract_common_headers(&self.token_extractor, headers);
116
117 let user_type = jwt_context.user_type.reconcile_with(user_is_admin(&user));
118 let session_id = jwt_context.session_id.clone();
119 let user_id = jwt_context.user_id.clone();
120
121 Ok(build_context(BuildContextParams {
122 jwt_context,
123 session_id,
124 user_id,
125 trace_id,
126 context_id,
127 agent_name,
128 task_id,
129 auth_token,
130 user_type,
131 }))
132 }
133
134 pub async fn decode_for_gateway(
135 &self,
136 jwt_token: &systemprompt_identifiers::JwtToken,
137 ) -> Result<(JwtUserContext, systemprompt_traits::AuthUser), ContextExtractionError> {
138 let jwt_context = extract_user_context(jwt_token.as_str(), &self.issuer)
139 .map_err(|e| ContextExtractionError::InvalidToken(e.into()))?;
140
141 let user = self.validate(&jwt_context, "gateway").await?;
142 Ok((jwt_context, user))
143 }
144
145 async fn extract_from_request_impl(
146 &self,
147 request: Request<Body>,
148 ) -> Result<(RequestContext, Request<Body>), ContextExtractionError> {
149 use crate::services::middleware::context::sources::{ContextIdSource, PayloadSource};
150
151 let headers = request.headers().clone();
152 let has_auth = headers.get("authorization").is_some();
153
154 if headers.get("x-context-id").is_some() && !has_auth {
155 return Err(ContextExtractionError::ForbiddenHeader {
156 header: "X-Context-ID".to_owned(),
157 reason: "Context ID must be in request body (A2A spec). Use contextId field in \
158 message."
159 .to_owned(),
160 });
161 }
162
163 let jwt_context = self.extract_jwt_context(&headers)?;
164 let user = self.validate(&jwt_context, " (A2A route)").await?;
165
166 let (body_bytes, reconstructed_request) =
167 PayloadSource::read_and_reconstruct(request).await?;
168
169 let context_source = PayloadSource::extract_context_source(&body_bytes)?;
170 let (context_id, task_id_from_payload) = match context_source {
171 ContextIdSource::Direct(id) => (
172 ContextId::try_new(id).map_err(|_invalid_id| {
173 ContextExtractionError::InvalidHeaderValue {
174 header: "contextId".to_owned(),
175 reason: "not a valid context id".to_owned(),
176 }
177 })?,
178 None,
179 ),
180 ContextIdSource::FromTask { task_id } => {
181 (ContextId::derived_from_task(&task_id), Some(task_id))
182 },
183 };
184
185 let (trace_id, task_id_from_header, auth_token, agent_name) =
186 extract_common_headers(&self.token_extractor, &headers);
187
188 let task_id = task_id_from_payload.or(task_id_from_header);
189 let user_type = jwt_context.user_type.reconcile_with(user_is_admin(&user));
190
191 let session_id = jwt_context.session_id.clone();
192 let user_id = jwt_context.user_id.clone();
193 let ctx = build_context(BuildContextParams {
194 jwt_context,
195 session_id,
196 user_id,
197 trace_id,
198 context_id,
199 agent_name,
200 task_id,
201 auth_token,
202 user_type,
203 });
204
205 Ok((ctx, reconstructed_request))
206 }
207}
208
209#[async_trait]
210impl ContextExtractor for JwtContextExtractor {
211 async fn extract_from_headers(
212 &self,
213 headers: &HeaderMap,
214 ) -> Result<RequestContext, ContextExtractionError> {
215 self.extract_standard(headers).await
216 }
217
218 async fn extract_from_request(
219 &self,
220 request: Request<Body>,
221 ) -> Result<(RequestContext, Request<Body>), ContextExtractionError> {
222 self.extract_from_request_impl(request).await
223 }
224}