systemprompt_security/extraction/
header.rs1use 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}