Skip to main content

systemprompt_security/extraction/
header.rs

1//! Header-based credential extraction and injection.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use axum::http::{HeaderMap, HeaderValue};
7use std::error::Error;
8use std::fmt;
9use systemprompt_identifiers::{
10    AgentName, ContextId, GatewayConversationId, ProviderRequestId, SessionId, TaskId, TraceId,
11    UserId, headers,
12};
13use systemprompt_models::execution::context::RequestContext;
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
16pub struct HeaderInjectionError;
17
18impl fmt::Display for HeaderInjectionError {
19    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
20        write!(f, "Header value contains invalid characters")
21    }
22}
23
24impl Error for HeaderInjectionError {}
25
26#[derive(Debug, Clone, Copy)]
27pub struct HeaderExtractor;
28
29impl HeaderExtractor {
30    pub fn extract_trace_id(headers: &HeaderMap) -> TraceId {
31        Self::extract_header(headers, headers::TRACE_ID)
32            .map_or_else(TraceId::generate, TraceId::new)
33    }
34
35    pub fn extract_context_id(headers: &HeaderMap) -> Option<ContextId> {
36        Self::extract_header(headers, headers::CONTEXT_ID)
37            .filter(|s| !s.is_empty())
38            .and_then(|s| {
39                ContextId::try_new(s)
40                    .map_err(|e| {
41                        tracing::warn!(error = %e, "Invalid context_id header value, ignoring");
42                        e
43                    })
44                    .ok()
45            })
46    }
47
48    pub fn extract_gateway_conversation_id(headers: &HeaderMap) -> Option<GatewayConversationId> {
49        Self::extract_header(headers, headers::GATEWAY_CONVERSATION_ID)
50            .filter(|s| !s.is_empty())
51            .and_then(|s| {
52                GatewayConversationId::try_new(s)
53                    .map_err(|e| {
54                        tracing::warn!(error = %e, "Invalid gateway_conversation_id header value, ignoring");
55                        e
56                    })
57                    .ok()
58            })
59    }
60
61    pub fn extract_provider_request_id(headers: &HeaderMap) -> Option<ProviderRequestId> {
62        Self::extract_header(headers, headers::PROVIDER_REQUEST_ID)
63            .filter(|s| !s.is_empty())
64            .and_then(|s| {
65                ProviderRequestId::try_new(s)
66                    .map_err(|e| {
67                        tracing::warn!(error = %e, "Invalid provider_request_id header value, ignoring");
68                        e
69                    })
70                    .ok()
71            })
72    }
73
74    pub fn extract_task_id(headers: &HeaderMap) -> Option<TaskId> {
75        Self::extract_header(headers, headers::TASK_ID).map(TaskId::new)
76    }
77
78    pub fn extract_agent_name(headers: &HeaderMap) -> AgentName {
79        Self::extract_header(headers, headers::AGENT_NAME)
80            .and_then(|s| {
81                AgentName::try_new(s)
82                    .map_err(|e| {
83                        tracing::warn!(error = %e, "Invalid agent_name header value, using system");
84                        e
85                    })
86                    .ok()
87            })
88            .unwrap_or_else(AgentName::system)
89    }
90
91    fn extract_header(headers: &HeaderMap, name: &str) -> Option<String> {
92        headers
93            .get(name)
94            .and_then(|v| {
95                v.to_str()
96                    .map_err(|e| {
97                        tracing::debug!(error = %e, header = %name, "Header contains non-ASCII characters");
98                        e
99                    })
100                    .ok()
101            })
102            .map(str::to_owned)
103    }
104}
105
106#[derive(Debug, Clone, Copy)]
107pub struct HeaderInjector;
108
109impl HeaderInjector {
110    pub fn inject_session_id(
111        headers: &mut HeaderMap,
112        session_id: &SessionId,
113    ) -> Result<(), HeaderInjectionError> {
114        Self::inject_header(headers, headers::SESSION_ID, session_id.as_str())
115    }
116
117    pub fn inject_user_id(
118        headers: &mut HeaderMap,
119        user_id: &UserId,
120    ) -> Result<(), HeaderInjectionError> {
121        Self::inject_header(headers, headers::USER_ID, user_id.as_str())
122    }
123
124    pub fn inject_trace_id(
125        headers: &mut HeaderMap,
126        trace_id: &TraceId,
127    ) -> Result<(), HeaderInjectionError> {
128        Self::inject_header(headers, headers::TRACE_ID, trace_id.as_str())
129    }
130
131    pub fn inject_context_id(
132        headers: &mut HeaderMap,
133        context_id: &ContextId,
134    ) -> Result<(), HeaderInjectionError> {
135        Self::inject_header(headers, headers::CONTEXT_ID, context_id.as_str())
136    }
137
138    pub fn inject_gateway_conversation_id(
139        headers: &mut HeaderMap,
140        id: &GatewayConversationId,
141    ) -> Result<(), HeaderInjectionError> {
142        Self::inject_header(headers, headers::GATEWAY_CONVERSATION_ID, id.as_str())
143    }
144
145    pub fn inject_provider_request_id(
146        headers: &mut HeaderMap,
147        id: &ProviderRequestId,
148    ) -> Result<(), HeaderInjectionError> {
149        Self::inject_header(headers, headers::PROVIDER_REQUEST_ID, id.as_str())
150    }
151
152    pub fn inject_task_id(
153        headers: &mut HeaderMap,
154        task_id: &TaskId,
155    ) -> Result<(), HeaderInjectionError> {
156        Self::inject_header(headers, headers::TASK_ID, task_id.as_str())
157    }
158
159    pub fn inject_agent_name(
160        headers: &mut HeaderMap,
161        agent_name: &str,
162    ) -> Result<(), HeaderInjectionError> {
163        Self::inject_header(headers, headers::AGENT_NAME, agent_name)
164    }
165
166    pub fn inject_from_request_context(
167        headers: &mut HeaderMap,
168        ctx: &RequestContext,
169    ) -> Result<(), HeaderInjectionError> {
170        Self::inject_session_id(headers, &ctx.request.session_id)?;
171        Self::inject_user_id(headers, &ctx.auth.actor.user_id)?;
172        Self::inject_trace_id(headers, &ctx.execution.trace_id)?;
173        Self::inject_context_id(headers, &ctx.execution.context_id)?;
174        Self::inject_agent_name(headers, ctx.execution.agent_name.as_str())?;
175        Ok(())
176    }
177
178    fn inject_header(
179        headers: &mut HeaderMap,
180        name: &'static str,
181        value: &str,
182    ) -> Result<(), HeaderInjectionError> {
183        HeaderValue::from_str(value).map_or(Err(HeaderInjectionError), |header_value| {
184            headers.insert(name, header_value);
185            Ok(())
186        })
187    }
188}