#![cfg_attr(test, allow(clippy::unwrap_used, clippy::expect_used))]
use std::cell::RefCell;
use pounce_common::types::{Index, Number};
use pounce_nlp::solve_statistics::IterRecord;
use tracing::field::{Field, Visit};
use tracing_subscriber::layer::{Context, Layer};
use tracing_subscriber::registry::LookupSpan;
pub const ITER_TARGET: &str = "pounce::iteration";
pub const RESTORATION_SPAN: &str = "restoration";
thread_local! {
static CAPTURE: RefCell<Option<Vec<IterRecord>>> = const { RefCell::new(None) };
}
static JSON_LOGGING: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
static GLOBAL_COLLECTOR: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
pub fn iteration_event_wanted() -> bool {
if JSON_LOGGING.load(std::sync::atomic::Ordering::Relaxed) {
return true;
}
CAPTURE.with(|c| c.borrow().is_some())
}
#[must_use = "call finish() to retrieve the captured iteration history"]
pub struct IterCaptureGuard {
prev: Option<Vec<IterRecord>>,
}
impl IterCaptureGuard {
pub fn start() -> Self {
let prev = CAPTURE.with(|c| c.borrow_mut().replace(Vec::new()));
Self { prev }
}
pub fn finish(mut self) -> Vec<IterRecord> {
let prev = self.prev.take();
let captured = CAPTURE
.with(|c| std::mem::replace(&mut *c.borrow_mut(), prev))
.unwrap_or_default();
std::mem::forget(self);
captured
}
}
impl Drop for IterCaptureGuard {
fn drop(&mut self) {
let prev = self.prev.take();
CAPTURE.with(|c| *c.borrow_mut() = prev);
}
}
fn push_record(rec: IterRecord) {
CAPTURE.with(|c| {
if let Some(buf) = c.borrow_mut().as_mut() {
buf.push(rec);
}
});
}
pub fn extend_active_capture(records: &[IterRecord]) {
if records.is_empty() {
return;
}
CAPTURE.with(|c| {
if let Some(buf) = c.borrow_mut().as_mut() {
buf.extend_from_slice(records);
}
});
}
#[derive(Default)]
struct IterVisitor {
rec: IterRecord,
}
impl Visit for IterVisitor {
fn record_f64(&mut self, field: &Field, value: f64) {
let v = value as Number;
match field.name() {
"objective" => self.rec.objective = v,
"inf_pr" => self.rec.inf_pr = v,
"inf_du" => self.rec.inf_du = v,
"mu" => self.rec.mu = v,
"d_norm" => self.rec.d_norm = v,
"regularization" => self.rec.regularization = v,
"alpha_dual" => self.rec.alpha_dual = v,
"alpha_primal" => self.rec.alpha_primal = v,
_ => {}
}
}
fn record_i64(&mut self, field: &Field, value: i64) {
match field.name() {
"iter" => self.rec.iter = value as Index,
"ls_trials" => self.rec.ls_trials = value as Index,
_ => {}
}
}
fn record_u64(&mut self, field: &Field, value: u64) {
match field.name() {
"iter" => self.rec.iter = value as Index,
"ls_trials" => self.rec.ls_trials = value as Index,
_ => {}
}
}
fn record_str(&mut self, field: &Field, value: &str) {
if field.name() == "alpha_char" {
self.rec.alpha_primal_char = value.chars().next().unwrap_or(' ');
}
}
fn record_debug(&mut self, _field: &Field, _value: &dyn std::fmt::Debug) {
}
}
#[derive(Debug, Default, Clone)]
pub struct IterCollectorLayer;
impl<S> Layer<S> for IterCollectorLayer
where
S: tracing::Subscriber + for<'a> LookupSpan<'a>,
{
fn on_event(&self, event: &tracing::Event<'_>, ctx: Context<'_, S>) {
if event.metadata().target() != ITER_TARGET {
return;
}
if let Some(scope) = ctx.event_scope(event) {
for span in scope.from_root() {
if span.name() == RESTORATION_SPAN {
return;
}
}
}
let mut visitor = IterVisitor::default();
event.record(&mut visitor);
push_record(visitor.rec);
}
}
fn collector_admits(m: &tracing::Metadata<'_>) -> bool {
m.is_span() || m.target() == ITER_TARGET
}
#[must_use = "the collector uninstalls as soon as this guard drops"]
pub struct CollectorScope {
_default: Option<tracing::subscriber::DefaultGuard>,
}
pub fn collector_scope() -> CollectorScope {
use tracing_subscriber::filter::filter_fn;
use tracing_subscriber::prelude::*;
if GLOBAL_COLLECTOR.load(std::sync::atomic::Ordering::Relaxed) {
return CollectorScope { _default: None };
}
let collector = IterCollectorLayer.with_filter(filter_fn(collector_admits));
let subscriber = tracing_subscriber::registry().with(collector);
CollectorScope {
_default: Some(tracing::subscriber::set_default(subscriber)),
}
}
#[must_use = "call finish() to retrieve the captured iteration history"]
pub struct ScopedIterCapture {
capture: IterCaptureGuard,
_scope: CollectorScope,
}
impl ScopedIterCapture {
pub fn start() -> Self {
let scope = collector_scope();
let capture = IterCaptureGuard::start();
Self {
capture,
_scope: scope,
}
}
pub fn finish(self) -> Vec<IterRecord> {
let Self { capture, _scope } = self;
let records = capture.finish();
drop(_scope);
records
}
}
pub fn with_iter_capture<R>(f: impl FnOnce() -> R) -> (R, Vec<IterRecord>) {
let scope = ScopedIterCapture::start();
let result = f();
(result, scope.finish())
}
fn level_style(level: tracing::Level) -> anstyle::Style {
use pounce_common::style::{ALPHA_HOT, TAN, TIGER_ORANGE};
let color = match level {
tracing::Level::ERROR => ALPHA_HOT,
tracing::Level::WARN => TIGER_ORANGE,
tracing::Level::INFO => TAN,
tracing::Level::DEBUG => anstyle::RgbColor(0x9a, 0x8c, 0x70),
tracing::Level::TRACE => anstyle::RgbColor(0x6a, 0x5d, 0x48),
};
anstyle::Style::new().fg_color(Some(anstyle::Color::Rgb(color)))
}
struct TigerFormat;
impl<S, N> tracing_subscriber::fmt::FormatEvent<S, N> for TigerFormat
where
S: tracing::Subscriber + for<'a> LookupSpan<'a>,
N: for<'a> tracing_subscriber::fmt::FormatFields<'a> + 'static,
{
fn format_event(
&self,
ctx: &tracing_subscriber::fmt::FmtContext<'_, S, N>,
mut writer: tracing_subscriber::fmt::format::Writer<'_>,
event: &tracing::Event<'_>,
) -> std::fmt::Result {
let meta = event.metadata();
let level = *meta.level();
if writer.has_ansi_escapes() {
let style = level_style(level);
write!(
writer,
"{}{:>5}{} ",
style.render(),
level,
style.render_reset()
)?;
} else {
write!(writer, "{level:>5} ")?;
}
write!(writer, "{}: ", meta.target())?;
ctx.field_format().format_fields(writer.by_ref(), event)?;
writeln!(writer)
}
}
pub fn init_subscriber() {
install();
}
pub fn init_for_tests() {
install();
}
fn install() {
use tracing_subscriber::EnvFilter;
use tracing_subscriber::filter::filter_fn;
use tracing_subscriber::prelude::*;
let _ = tracing_log::LogTracer::init();
let want_json = std::env::var("POUNCE_LOG_FORMAT")
.map(|v| v.eq_ignore_ascii_case("json"))
.unwrap_or(false);
JSON_LOGGING.store(want_json, std::sync::atomic::Ordering::Relaxed);
let claimed = if want_json {
let collector = IterCollectorLayer.with_filter(filter_fn(collector_admits));
let json_layer = tracing_subscriber::fmt::layer()
.json()
.with_writer(std::io::stderr)
.with_filter(env_filter());
tracing_subscriber::registry()
.with(json_layer)
.with(collector)
.try_init()
.is_ok()
} else {
let collector = IterCollectorLayer.with_filter(filter_fn(collector_admits));
let ansi = ansi_enabled();
let text_layer = tracing_subscriber::fmt::layer()
.event_format(TigerFormat)
.with_ansi(ansi)
.with_writer(std::io::stderr)
.with_filter(console_filter());
tracing_subscriber::registry()
.with(text_layer)
.with(collector)
.try_init()
.is_ok()
};
if claimed {
GLOBAL_COLLECTOR.store(true, std::sync::atomic::Ordering::Relaxed);
}
fn env_filter() -> EnvFilter {
EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info"))
}
fn console_filter() -> EnvFilter {
let base = env_filter();
match format!("{ITER_TARGET}=off").parse() {
Ok(directive) => base.add_directive(directive),
Err(_) => base,
}
}
fn ansi_enabled() -> bool {
if anstyle_query::clicolor_force() {
return true;
}
if anstyle_query::no_color() {
return false;
}
anstyle_query::term_supports_ansi_color()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_record(iter: i32, alpha: f64, c: char) -> IterRecord {
IterRecord {
iter,
objective: 1.0,
inf_pr: 2.0,
inf_du: 3.0,
mu: 4.0,
d_norm: 5.0,
regularization: 6.0,
alpha_dual: 7.0,
alpha_primal: alpha,
alpha_primal_char: c,
ls_trials: 1,
}
}
#[test]
fn iteration_event_wanted_tracks_active_capture() {
assert!(!iteration_event_wanted());
let guard = IterCaptureGuard::start();
assert!(iteration_event_wanted(), "capture active → event wanted");
let _ = guard.finish();
assert!(
!iteration_event_wanted(),
"capture ended → event suppressed"
);
}
#[test]
fn guard_captures_pushed_records() {
let guard = IterCaptureGuard::start();
push_record(sample_record(0, 1.0, ' '));
push_record(sample_record(1, 0.5, 'R'));
let got = guard.finish();
assert_eq!(got.len(), 2);
assert_eq!(got[1].iter, 1);
assert_eq!(got[1].alpha_primal_char, 'R');
}
#[test]
fn no_guard_means_records_are_dropped() {
push_record(sample_record(0, 1.0, ' '));
let guard = IterCaptureGuard::start();
let got = guard.finish();
assert!(got.is_empty());
}
#[test]
fn guard_restores_previous_slot_on_finish() {
let outer = IterCaptureGuard::start();
push_record(sample_record(0, 1.0, ' '));
{
let inner = IterCaptureGuard::start();
push_record(sample_record(99, 0.1, 'R'));
let inner_got = inner.finish();
assert_eq!(inner_got.len(), 1);
assert_eq!(inner_got[0].iter, 99);
}
push_record(sample_record(1, 1.0, ' '));
let outer_got = outer.finish();
assert_eq!(outer_got.len(), 2);
assert_eq!(outer_got[0].iter, 0);
assert_eq!(outer_got[1].iter, 1);
}
#[test]
fn collector_excludes_restoration_nested_iterations() {
use tracing_subscriber::filter::filter_fn;
use tracing_subscriber::prelude::*;
fn emit(iter: i64, ch: char) {
let s = ch.to_string();
tracing::info!(
target: ITER_TARGET,
iter = iter,
objective = 0.0,
alpha_primal = 1.0,
alpha_char = s.as_str(),
);
}
let collector = IterCollectorLayer.with_filter(filter_fn(collector_admits));
let subscriber = tracing_subscriber::registry().with(collector);
let captured = tracing::subscriber::with_default(subscriber, || {
let guard = IterCaptureGuard::start();
emit(0, ' '); {
let _resto = tracing::info_span!("restoration").entered();
let _inner_solve = tracing::info_span!("solve").entered();
let _inner_iter = tracing::info_span!("iteration").entered();
emit(99, 'R'); }
emit(1, ' '); guard.finish()
});
let iters: Vec<i32> = captured.iter().map(|r| r.iter).collect();
assert_eq!(
iters,
vec![0, 1],
"inner restoration iteration leaked: {iters:?}"
);
}
#[test]
fn log_records_bridge_into_tracing() {
use std::sync::{Arc, Mutex};
use tracing_subscriber::prelude::*;
#[derive(Clone)]
struct CaptureLayer {
buf: Arc<Mutex<Vec<String>>>,
}
impl<S: tracing::Subscriber> tracing_subscriber::Layer<S> for CaptureLayer {
fn on_event(
&self,
event: &tracing::Event<'_>,
_ctx: tracing_subscriber::layer::Context<'_, S>,
) {
struct V<'a>(&'a mut Vec<String>);
impl tracing::field::Visit for V<'_> {
fn record_debug(&mut self, f: &Field, value: &dyn std::fmt::Debug) {
if f.name() == "message" {
self.0.push(format!("{value:?}"));
}
}
}
let mut g = self.buf.lock().unwrap_or_else(|p| p.into_inner());
event.record(&mut V(&mut g));
}
}
let buf = Arc::new(Mutex::new(Vec::new()));
let subscriber = tracing_subscriber::registry().with(CaptureLayer { buf: buf.clone() });
let _ = tracing_log::LogTracer::init();
tracing::subscriber::with_default(subscriber, || {
log::error!(target: "some_transitive_dep", "bridged log record");
});
let got = buf.lock().unwrap_or_else(|p| p.into_inner());
assert!(
got.iter().any(|m| m.contains("bridged log record")),
"log record did not reach the tracing layer; captured: {got:?}"
);
}
fn emit_iter(iter: i64, ch: char) {
let s = ch.to_string();
tracing::info!(
target: ITER_TARGET,
iter = iter,
objective = 0.5,
alpha_primal = 1.0,
alpha_char = s.as_str(),
);
}
#[test]
fn extend_active_capture_appends_to_enclosing_buffer() {
let outer = IterCaptureGuard::start();
push_record(sample_record(0, 1.0, ' '));
let inner = IterCaptureGuard::start();
push_record(sample_record(1, 0.5, ' '));
let inner_got = inner.finish();
extend_active_capture(&inner_got);
let outer_got = outer.finish();
let iters: Vec<i32> = outer_got.iter().map(|r| r.iter).collect();
assert_eq!(iters, vec![0, 1]);
extend_active_capture(&inner_got);
let fresh = IterCaptureGuard::start();
assert!(fresh.finish().is_empty());
}
#[test]
fn with_iter_capture_captures_events_and_threads_result() {
let (result, records) = with_iter_capture(|| {
emit_iter(0, ' ');
emit_iter(1, 'R');
"sentinel"
});
assert_eq!(result, "sentinel");
assert_eq!(records.len(), 2);
assert_eq!(records[0].iter, 0);
assert_eq!(records[1].iter, 1);
assert_eq!(records[1].alpha_primal_char, 'R');
assert!((records[1].objective - 0.5).abs() < 1e-12);
}
#[test]
fn with_iter_capture_ignores_events_outside_scope() {
emit_iter(7, ' '); let ((), records) = with_iter_capture(|| ());
emit_iter(8, ' '); assert!(records.is_empty());
assert!(
!iteration_event_wanted(),
"capture slot must be torn down after with_iter_capture returns"
);
}
#[test]
fn with_iter_capture_excludes_restoration_subsolve() {
let ((), records) = with_iter_capture(|| {
emit_iter(0, ' ');
{
let _resto = tracing::info_span!("restoration").entered();
let _inner_iter = tracing::info_span!("iteration").entered();
emit_iter(99, 'R');
}
emit_iter(1, ' ');
});
let iters: Vec<i32> = records.iter().map(|r| r.iter).collect();
assert_eq!(
iters,
vec![0, 1],
"inner restoration iteration leaked: {iters:?}"
);
}
#[test]
fn scoped_iter_capture_nesting_restores_outer_buffer() {
let outer = ScopedIterCapture::start();
emit_iter(0, ' ');
let ((), inner) = with_iter_capture(|| emit_iter(99, 'R'));
emit_iter(1, ' ');
let outer_got = outer.finish();
assert_eq!(inner.len(), 1);
assert_eq!(inner[0].iter, 99);
let iters: Vec<i32> = outer_got.iter().map(|r| r.iter).collect();
assert_eq!(iters, vec![0, 1]);
}
#[test]
fn collector_scope_feeds_manual_guard() {
let scope = collector_scope();
let guard = IterCaptureGuard::start();
emit_iter(0, ' ');
let records = guard.finish();
drop(scope);
assert_eq!(records.len(), 1);
assert_eq!(records[0].iter, 0);
let guard = IterCaptureGuard::start();
emit_iter(1, ' ');
assert!(guard.finish().is_empty());
}
#[test]
fn with_iter_capture_is_panic_safe() {
let unwound = std::panic::catch_unwind(|| {
let _ = with_iter_capture(|| panic!("solve blew up"));
});
assert!(unwound.is_err());
assert!(
!iteration_event_wanted(),
"capture slot must be restored when the closure unwinds"
);
}
#[test]
fn iter_record_default_and_assignment() {
let mut v = IterVisitor::default();
v.rec.iter = 7;
v.rec.alpha_primal = 0.25;
v.rec.alpha_primal_char = 'S';
assert_eq!(v.rec.iter, 7);
assert_eq!(v.rec.alpha_primal_char, 'S');
}
}