Skip to main content

meow_api/
log_stream.rs

1use serde::ser::{SerializeStruct, Serializer};
2use serde::Serialize;
3use tokio::sync::broadcast;
4use tracing_subscriber::Layer;
5
6type LogReload = dyn Fn(&str) -> Result<(), String> + Send + Sync;
7static LOG_RELOAD: std::sync::OnceLock<Box<LogReload>> = std::sync::OnceLock::new();
8
9pub fn install_log_reloader<F>(reload: F)
10where
11    F: Fn(&str) -> Result<(), String> + Send + Sync + 'static,
12{
13    let _ = LOG_RELOAD.set(Box::new(reload));
14}
15
16pub fn reload_log_level(level: &str) -> Result<(), String> {
17    match LOG_RELOAD.get() {
18        Some(reload) => reload(level),
19        None => Ok(()),
20    }
21}
22
23#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
24pub enum LogLevel {
25    Debug,
26    Info,
27    Warning,
28    Error,
29    Silent,
30}
31
32impl LogLevel {
33    pub fn as_str(self) -> &'static str {
34        match self {
35            LogLevel::Debug => "debug",
36            LogLevel::Info => "info",
37            LogLevel::Warning => "warning",
38            LogLevel::Error => "error",
39            LogLevel::Silent => "silent",
40        }
41    }
42}
43
44#[derive(Clone, Debug)]
45pub struct LogMessage {
46    pub level: LogLevel,
47    pub payload: String,
48    pub time: time::OffsetDateTime,
49}
50
51impl Serialize for LogMessage {
52    fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
53        let mut m = s.serialize_struct("LogMessage", 3)?;
54        m.serialize_field("type", self.level.as_str())?;
55        m.serialize_field("payload", &self.payload)?;
56        let ts = self
57            .time
58            .format(&time::format_description::well_known::Rfc3339)
59            .unwrap_or_default();
60        m.serialize_field("time", &ts)?;
61        m.end()
62    }
63}
64
65pub struct LogBroadcastLayer {
66    pub tx: broadcast::Sender<LogMessage>,
67}
68
69impl<S: tracing::Subscriber> Layer<S> for LogBroadcastLayer {
70    fn on_event(
71        &self,
72        event: &tracing::Event<'_>,
73        _ctx: tracing_subscriber::layer::Context<'_, S>,
74    ) {
75        let level = match *event.metadata().level() {
76            tracing::Level::TRACE | tracing::Level::DEBUG => LogLevel::Debug,
77            tracing::Level::INFO => LogLevel::Info,
78            tracing::Level::WARN => LogLevel::Warning,
79            tracing::Level::ERROR => LogLevel::Error,
80        };
81        let mut visitor = MessageVisitor(String::new());
82        event.record(&mut visitor);
83        let msg = LogMessage {
84            level,
85            payload: visitor.0,
86            time: time::OffsetDateTime::now_utc(),
87        };
88        // Non-blocking; Err = no subscribers or channel full — both acceptable.
89        let _ = self.tx.send(msg);
90    }
91}
92
93struct MessageVisitor(String);
94
95impl tracing::field::Visit for MessageVisitor {
96    fn record_debug(&mut self, field: &tracing::field::Field, value: &dyn std::fmt::Debug) {
97        if field.name() == "message" {
98            self.0 = format!("{value:?}");
99        }
100    }
101
102    fn record_str(&mut self, field: &tracing::field::Field, value: &str) {
103        if field.name() == "message" {
104            self.0 = value.to_string();
105        }
106    }
107}
108
109pub fn parse_log_level(s: &str) -> LogLevel {
110    match s.to_ascii_lowercase().as_str() {
111        "debug" => LogLevel::Debug,
112        "warning" | "warn" => LogLevel::Warning,
113        "error" => LogLevel::Error,
114        "silent" => LogLevel::Silent,
115        // "info" and any unrecognised value default to Info.
116        _ => LogLevel::Info,
117    }
118}
119
120#[cfg(test)]
121mod tests {
122    use super::*;
123    use tracing::subscriber;
124    use tracing_subscriber::prelude::*;
125
126    #[test]
127    fn log_level_str_round_trip_via_parse() {
128        for (s, lvl) in [
129            ("debug", LogLevel::Debug),
130            ("info", LogLevel::Info),
131            ("warning", LogLevel::Warning),
132            ("error", LogLevel::Error),
133            ("silent", LogLevel::Silent),
134        ] {
135            assert_eq!(parse_log_level(s), lvl);
136            assert_eq!(lvl.as_str(), s);
137        }
138    }
139
140    #[test]
141    fn parse_log_level_accepts_warn_alias() {
142        assert_eq!(parse_log_level("warn"), LogLevel::Warning);
143        assert_eq!(parse_log_level("WARN"), LogLevel::Warning);
144    }
145
146    #[test]
147    fn parse_log_level_is_case_insensitive() {
148        assert_eq!(parse_log_level("DEBUG"), LogLevel::Debug);
149        assert_eq!(parse_log_level("Info"), LogLevel::Info);
150        assert_eq!(parse_log_level("Error"), LogLevel::Error);
151    }
152
153    #[test]
154    fn parse_log_level_unknown_input_defaults_to_info() {
155        // Documented behaviour: an unrecognised level → fall back to Info
156        // rather than rejecting the request.
157        assert_eq!(parse_log_level(""), LogLevel::Info);
158        assert_eq!(parse_log_level("nonsense"), LogLevel::Info);
159        assert_eq!(parse_log_level("trace"), LogLevel::Info);
160    }
161
162    #[test]
163    fn log_level_ord_matches_severity_increasing() {
164        // The WS handler filters with `level >= request_level`. The enum
165        // ordering must therefore be Debug < Info < Warning < Error < Silent
166        // (where Silent is the strictest filter — nothing passes).
167        assert!(LogLevel::Debug < LogLevel::Info);
168        assert!(LogLevel::Info < LogLevel::Warning);
169        assert!(LogLevel::Warning < LogLevel::Error);
170        assert!(LogLevel::Error < LogLevel::Silent);
171    }
172
173    #[test]
174    fn log_message_serializes_three_fields() {
175        // 2026-01-02 03:04:05 UTC = 1767322945 seconds since unix epoch.
176        let ts = time::OffsetDateTime::from_unix_timestamp(1_767_322_945).unwrap();
177        let msg = LogMessage {
178            level: LogLevel::Warning,
179            payload: "alerts: thing happened".into(),
180            time: ts,
181        };
182        let json = serde_json::to_value(&msg).unwrap();
183        assert_eq!(json["type"], "warning");
184        assert_eq!(json["payload"], "alerts: thing happened");
185        assert_eq!(json["time"], "2026-01-02T03:02:25Z");
186    }
187
188    #[test]
189    fn layer_forwards_string_message_event() {
190        let (tx, mut rx) = broadcast::channel(8);
191        let layer = LogBroadcastLayer { tx };
192        let registry = tracing_subscriber::registry().with(layer);
193        subscriber::with_default(registry, || {
194            tracing::warn!("hello-world");
195        });
196        let got = rx.try_recv().expect("event must be forwarded");
197        assert_eq!(got.level, LogLevel::Warning);
198        assert!(
199            got.payload.contains("hello-world"),
200            "payload: {}",
201            got.payload
202        );
203    }
204
205    #[test]
206    fn layer_maps_tracing_levels_to_log_levels() {
207        let (tx, mut rx) = broadcast::channel(16);
208        let layer = LogBroadcastLayer { tx };
209        let registry = tracing_subscriber::registry().with(layer);
210        subscriber::with_default(registry, || {
211            tracing::error!("e");
212            tracing::warn!("w");
213            tracing::info!("i");
214            tracing::debug!("d"); // collapses to Debug
215            tracing::trace!("t"); // collapses to Debug
216        });
217        let levels: Vec<LogLevel> = std::iter::from_fn(|| rx.try_recv().ok())
218            .map(|m| m.level)
219            .collect();
220        // Trace+debug both map to Debug; default filter may drop them at the
221        // subscriber level, so accept either {Error, Warning, Info} only or
222        // the full set.
223        assert!(levels.starts_with(&[LogLevel::Error, LogLevel::Warning, LogLevel::Info]));
224        for extra in &levels[3..] {
225            assert_eq!(*extra, LogLevel::Debug);
226        }
227    }
228
229    #[test]
230    fn layer_send_with_no_subscribers_does_not_panic() {
231        // Documented contract: a send Err (no subscribers / channel full) is
232        // acceptable — verify we don't regress to a panicking `.unwrap()`.
233        let (tx, rx) = broadcast::channel(1);
234        drop(rx);
235        let layer = LogBroadcastLayer { tx };
236        let registry = tracing_subscriber::registry().with(layer);
237        subscriber::with_default(registry, || {
238            tracing::info!("payload");
239            tracing::error!("oh no");
240        });
241        // No panic = pass.
242    }
243}