Skip to main content

systemprompt_logging/layer/
proxy.rs

1//! Deferred database-logging tracing layer.
2//!
3//! [`ProxyDatabaseLayer`] is installed in the subscriber stack before a
4//! database pool exists, buffering span attribution into span extensions. Once
5//! [`ProxyDatabaseLayer::attach`] supplies a pool it delegates to the real
6//! `DatabaseLayer`; until then span fields are recorded so attribution is not
7//! lost across the boot window. The free functions build the [`LogEntry`]
8//! actor triple by walking the span tree.
9
10use 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    /// Attaches the database sink, idempotently: the first pool wins and repeat
51    /// attaches are ignored. Repeats are expected — `init_logging` is reached
52    /// from more than one entry point during a single startup — so
53    /// `get_or_init` keeps this silent and avoids spawning a second writer
54    /// task.
55    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}