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) {
54 self.inner.get_or_init(|| DatabaseLayer::new(db_pool));
55 }
56}
57
58impl<S> Layer<S> for ProxyDatabaseLayer
59where
60 S: Subscriber + for<'a> LookupSpan<'a>,
61{
62 fn on_new_span(
63 &self,
64 attrs: &tracing::span::Attributes<'_>,
65 id: &tracing::span::Id,
66 ctx: Context<'_, S>,
67 ) {
68 if let Some(db) = self.inner.get() {
69 db.on_new_span(attrs, id, ctx);
70 } else {
71 record_span_fields(attrs, id, &ctx);
72 }
73 }
74
75 fn on_record(
76 &self,
77 id: &tracing::span::Id,
78 values: &tracing::span::Record<'_>,
79 ctx: Context<'_, S>,
80 ) {
81 if let Some(db) = self.inner.get() {
82 db.on_record(id, values, ctx);
83 } else {
84 update_span_fields(id, values, &ctx);
85 }
86 }
87
88 fn on_event(&self, event: &Event<'_>, ctx: Context<'_, S>) {
89 if let Some(db) = self.inner.get() {
90 db.on_event(event, ctx);
91 }
92 }
93}
94
95pub(super) fn record_span_fields<S>(
96 attrs: &tracing::span::Attributes<'_>,
97 id: &tracing::span::Id,
98 ctx: &Context<'_, S>,
99) where
100 S: Subscriber + for<'a> LookupSpan<'a>,
101{
102 let Some(span) = ctx.span(id) else {
103 return;
104 };
105 let mut fields = SpanFields::default();
106 let mut context = SpanContext::default();
107 let mut visitor = SpanVisitor {
108 context: &mut context,
109 };
110 attrs.record(&mut visitor);
111
112 fields.user = context.user;
113 fields.session = context.session;
114 fields.task = context.task;
115 fields.trace = context.trace;
116 fields.context = context.context;
117 fields.client = context.client;
118
119 let mut extensions = span.extensions_mut();
120 extensions.insert(fields);
121}
122
123pub(super) fn update_span_fields<S>(
124 id: &tracing::span::Id,
125 values: &tracing::span::Record<'_>,
126 ctx: &Context<'_, S>,
127) where
128 S: Subscriber + for<'a> LookupSpan<'a>,
129{
130 if let Some(span) = ctx.span(id) {
131 let mut extensions = span.extensions_mut();
132 if let Some(fields) = extensions.get_mut::<SpanFields>() {
133 let mut context = SpanContext {
134 user: fields.user.clone(),
135 session: fields.session.clone(),
136 task: fields.task.clone(),
137 trace: fields.trace.clone(),
138 context: fields.context.clone(),
139 client: fields.client.clone(),
140 };
141 let mut visitor = SpanVisitor {
142 context: &mut context,
143 };
144 values.record(&mut visitor);
145
146 fields.user = context.user;
147 fields.session = context.session;
148 fields.task = context.task;
149 fields.trace = context.trace;
150 fields.context = context.context;
151 fields.client = context.client;
152 }
153 }
154}
155
156pub(super) fn build_log_entry<S>(event: &Event<'_>, ctx: &Context<'_, S>) -> Option<LogEntry>
157where
158 S: Subscriber + for<'a> LookupSpan<'a>,
159{
160 let level = *event.metadata().level();
161 let module = event.metadata().target().to_owned();
162
163 let mut visitor = FieldVisitor::default();
164 event.record(&mut visitor);
165
166 let span_context = ctx
167 .current_span()
168 .id()
169 .and_then(|id| ctx.span(id))
170 .map(extract_span_context)?;
171
172 let log_level = match level {
173 tracing::Level::ERROR => LogLevel::Error,
174 tracing::Level::WARN => LogLevel::Warn,
175 tracing::Level::INFO => LogLevel::Info,
176 tracing::Level::DEBUG => LogLevel::Debug,
177 tracing::Level::TRACE => LogLevel::Trace,
178 };
179
180 let user_id = UserId::new(span_context.user.as_ref()?.clone());
181 let session_id = SessionId::new(span_context.session.as_ref()?.clone());
182 let trace_id = TraceId::new(span_context.trace.as_ref()?.clone());
183
184 Some(LogEntry {
185 id: LogId::generate(),
186 timestamp: Utc::now(),
187 level: log_level,
188 module,
189 message: visitor.message,
190 metadata: visitor.fields,
191 user_id,
192 session_id,
193 task_id: span_context.task.as_ref().map(|s| TaskId::new(s.clone())),
194 trace_id,
195 context_id: span_context.context.as_ref().and_then(|s| {
196 ContextId::try_new(s.clone())
197 .map_err(|e| {
198 tracing::warn!(error = %e, raw = %s, "Skipping non-UUID context_id from span context");
199 e
200 })
201 .ok()
202 }),
203 client_id: span_context
204 .client
205 .as_ref()
206 .map(|s| ClientId::new(s.clone())),
207 })
208}