Skip to main content

systemprompt_api/services/middleware/jwt/
params.rs

1//! JWT middleware construction parameters.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use axum::http::HeaderMap;
7use systemprompt_identifiers::{AgentName, ContextId, SessionId, TaskId, TraceId, UserId};
8use systemprompt_models::auth::UserType;
9use systemprompt_models::execution::context::RequestContext;
10use systemprompt_security::{HeaderExtractor, JwtUserContext, TokenExtractor};
11
12#[derive(Debug)]
13pub struct BuildContextParams {
14    pub jwt_context: JwtUserContext,
15    pub session_id: SessionId,
16    pub user_id: UserId,
17    pub trace_id: TraceId,
18    pub context_id: ContextId,
19    pub agent_name: AgentName,
20    pub task_id: Option<TaskId>,
21    pub auth_token: Option<String>,
22    pub user_type: UserType,
23}
24
25pub fn build_context(params: BuildContextParams) -> RequestContext {
26    let BuildContextParams {
27        jwt_context,
28        session_id,
29        user_id,
30        trace_id,
31        context_id,
32        agent_name,
33        task_id,
34        auth_token,
35        user_type,
36    } = params;
37    let mut ctx = RequestContext::new(session_id, trace_id, context_id, agent_name)
38        .with_actor(systemprompt_identifiers::Actor::user(user_id))
39        .with_user_type(user_type)
40        .with_act_chain(jwt_context.act_chain)
41        .with_jti(jwt_context.jti)
42        .with_token_exp(jwt_context.exp);
43
44    if let Some(client_id) = jwt_context.client_id {
45        ctx = ctx.with_client_id(client_id);
46    }
47    if let Some(t_id) = task_id {
48        ctx = ctx.with_task_id(t_id);
49    }
50    if let Some(token) = auth_token {
51        ctx = ctx.with_auth_token(token);
52    }
53    ctx
54}
55
56pub fn extract_common_headers(
57    token_extractor: &TokenExtractor,
58    headers: &HeaderMap,
59) -> (TraceId, Option<TaskId>, Option<String>, AgentName) {
60    (
61        HeaderExtractor::extract_trace_id(headers),
62        HeaderExtractor::extract_task_id(headers),
63        token_extractor.extract(headers).ok(),
64        HeaderExtractor::extract_agent_name(headers),
65    )
66}