systemprompt_logging/layer/
proxy.rs1use std::sync::{Arc, OnceLock};
11
12use chrono::Utc;
13use tracing::{Event, Subscriber};
14use tracing_subscriber::Layer;
15use tracing_subscriber::layer::Context;
16use tracing_subscriber::registry::LookupSpan;
17
18use super::DatabaseLayer;
19use super::visitor::{FieldVisitor, SpanContext, SpanFields, SpanVisitor, extract_span_context};
20use crate::models::{LogEntry, LogLevel};
21use systemprompt_database::DbPool;
22use systemprompt_identifiers::{ClientId, ContextId, LogId, SessionId, TaskId, TraceId, UserId};
23
24#[derive(Clone)]
25pub struct ProxyDatabaseLayer {
26 inner: Arc<OnceLock<DatabaseLayer>>,
27}
28
29impl std::fmt::Debug for ProxyDatabaseLayer {
30 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31 f.debug_struct("ProxyDatabaseLayer")
32 .field("attached", &self.inner.get().is_some())
33 .finish()
34 }
35}
36
37impl Default for ProxyDatabaseLayer {
38 fn default() -> Self {
39 Self::new()
40 }
41}
42
43impl ProxyDatabaseLayer {
44 pub fn new() -> Self {
45 Self {
46 inner: Arc::new(OnceLock::new()),
47 }
48 }
49
50 pub fn attach(&self, db_pool: DbPool) {
56 self.inner.get_or_init(|| DatabaseLayer::new(db_pool));
57 }
58}
59
60impl<S> Layer<S> for ProxyDatabaseLayer
61where
62 S: Subscriber + for<'a> LookupSpan<'a>,
63{
64 fn on_new_span(
65 &self,
66 attrs: &tracing::span::Attributes<'_>,
67 id: &tracing::span::Id,
68 ctx: Context<'_, S>,
69 ) {
70 if let Some(db) = self.inner.get() {
71 db.on_new_span(attrs, id, ctx);
72 } else {
73 record_span_fields(attrs, id, &ctx);
74 }
75 }
76
77 fn on_record(
78 &self,
79 id: &tracing::span::Id,
80 values: &tracing::span::Record<'_>,
81 ctx: Context<'_, S>,
82 ) {
83 if let Some(db) = self.inner.get() {
84 db.on_record(id, values, ctx);
85 } else {
86 update_span_fields(id, values, &ctx);
87 }
88 }
89
90 fn on_event(&self, event: &Event<'_>, ctx: Context<'_, S>) {
91 if let Some(db) = self.inner.get() {
92 db.on_event(event, ctx);
93 }
94 }
95}
96
97pub(super) fn record_span_fields<S>(
98 attrs: &tracing::span::Attributes<'_>,
99 id: &tracing::span::Id,
100 ctx: &Context<'_, S>,
101) where
102 S: Subscriber + for<'a> LookupSpan<'a>,
103{
104 let Some(span) = ctx.span(id) else {
105 return;
106 };
107 let mut fields = SpanFields::default();
108 let mut context = SpanContext::default();
109 let mut visitor = SpanVisitor {
110 context: &mut context,
111 };
112 attrs.record(&mut visitor);
113
114 fields.user = context.user;
115 fields.session = context.session;
116 fields.task = context.task;
117 fields.trace = context.trace;
118 fields.context = context.context;
119 fields.client = context.client;
120
121 let mut extensions = span.extensions_mut();
122 extensions.insert(fields);
123}
124
125pub(super) fn update_span_fields<S>(
126 id: &tracing::span::Id,
127 values: &tracing::span::Record<'_>,
128 ctx: &Context<'_, S>,
129) where
130 S: Subscriber + for<'a> LookupSpan<'a>,
131{
132 if let Some(span) = ctx.span(id) {
133 let mut extensions = span.extensions_mut();
134 if let Some(fields) = extensions.get_mut::<SpanFields>() {
135 let mut context = SpanContext {
136 user: fields.user.clone(),
137 session: fields.session.clone(),
138 task: fields.task.clone(),
139 trace: fields.trace.clone(),
140 context: fields.context.clone(),
141 client: fields.client.clone(),
142 };
143 let mut visitor = SpanVisitor {
144 context: &mut context,
145 };
146 values.record(&mut visitor);
147
148 fields.user = context.user;
149 fields.session = context.session;
150 fields.task = context.task;
151 fields.trace = context.trace;
152 fields.context = context.context;
153 fields.client = context.client;
154 }
155 }
156}
157
158pub(super) fn build_log_entry<S>(event: &Event<'_>, ctx: &Context<'_, S>) -> Option<LogEntry>
159where
160 S: Subscriber + for<'a> LookupSpan<'a>,
161{
162 let level = *event.metadata().level();
163 let module = event.metadata().target().to_owned();
164
165 let mut visitor = FieldVisitor::default();
166 event.record(&mut visitor);
167
168 let span_context = ctx
169 .current_span()
170 .id()
171 .and_then(|id| ctx.span(id))
172 .map(extract_span_context)?;
173
174 let log_level = match level {
175 tracing::Level::ERROR => LogLevel::Error,
176 tracing::Level::WARN => LogLevel::Warn,
177 tracing::Level::INFO => LogLevel::Info,
178 tracing::Level::DEBUG => LogLevel::Debug,
179 tracing::Level::TRACE => LogLevel::Trace,
180 };
181
182 let user_id = UserId::new(span_context.user.as_ref()?.clone());
183 let session_id = SessionId::new(span_context.session.as_ref()?.clone());
184 let trace_id = TraceId::new(span_context.trace.as_ref()?.clone());
185
186 Some(LogEntry {
187 id: LogId::generate(),
188 timestamp: Utc::now(),
189 level: log_level,
190 module,
191 message: visitor.message,
192 metadata: visitor.fields,
193 user_id,
194 session_id,
195 task_id: span_context.task.as_ref().map(|s| TaskId::new(s.clone())),
196 trace_id,
197 context_id: span_context.context.as_ref().and_then(|s| {
198 ContextId::try_new(s.clone())
199 .map_err(|e| {
200 tracing::warn!(error = %e, raw = %s, "Skipping non-UUID context_id from span context");
201 e
202 })
203 .ok()
204 }),
205 client_id: span_context
206 .client
207 .as_ref()
208 .map(|s| ClientId::new(s.clone())),
209 })
210}