Skip to main content

systemprompt_api/services/middleware/jwt/
context.rs

1//! JWT-backed request-context extractor.
2//!
3//! [`JwtContextExtractor`] implements [`ContextExtractor`] by validating the
4//! bearer token (signature, session existence, user existence, and JTI
5//! revocation) and building a `RequestContext`. It resolves the context id from
6//! the `x-context-id` header on standard routes and from the JSON-RPC body on
7//! A2A routes, and exposes a gateway decode path for pre-authenticated tokens.
8//!
9//! Copyright (c) systemprompt.io — Business Source License 1.1.
10//! See <https://systemprompt.io> for licensing details.
11
12use 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}