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