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//!
10//! Copyright (c) systemprompt.io — Business Source License 1.1.
11//! See <https://systemprompt.io> for licensing details.
12
13use 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    /// Attaches the database sink, idempotently: the first pool wins and repeat
54    /// attaches are ignored. Repeats are expected — `init_logging` is reached
55    /// from more than one entry point during a single startup — so
56    /// `get_or_init` keeps this silent and avoids spawning a second writer
57    /// task.
58    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}