use super::struct_test::{
FormatResult, Issue, QueryTrace, Reason, TestFailure, TraceSink, msg, msgf,
};
use super::transaction_test::TestTransaction;
use crate::utils::aliases::StrMap;
use futures_util::FutureExt;
use std::any::Any;
use std::panic::AssertUnwindSafe;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex};
use std::time::Duration;
static SERIAL: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
static LAST_TARGET: Mutex<Option<String>> = Mutex::new(None);
pub const TIMEOUT_KEY: &str = "RUNIQUE_TEST_TIMEOUT";
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
const CLEANUP_TIMEOUT: Duration = Duration::from_secs(15);
pub(crate) struct TestScope {
trace: TraceSink,
}
tokio::task_local! {
static CURRENT: Arc<TestScope>;
}
pub async fn runique_test<C: TestTransaction>(
env_file: &str,
handler: impl AsyncFnOnce(&C) -> Result<(), C::Error>,
) -> Result<(), TestFailure> {
let name_test = current_test_name();
if CURRENT.try_with(|_| ()).is_ok() {
return report(FormatResult::<C::Error> {
name_test,
reason: Reason::Setup(msg("runique_test.nested").into_owned()),
trace: Vec::new(),
issues: Vec::new(),
});
}
let _serial = SERIAL.lock().await;
let trace = TraceSink::default();
let scope = Arc::new(TestScope {
trace: trace.clone(),
});
let (reason, issues) = CURRENT
.scope(scope.clone(), run::<C>(env_file, &trace, handler))
.await;
report(FormatResult {
name_test,
reason,
trace: std::mem::take(&mut *lock(&trace)),
issues,
})
}
fn report<E: std::fmt::Display>(result: FormatResult<E>) -> Result<(), TestFailure> {
print!("{result}");
if result.passed() {
return Ok(());
}
Err(TestFailure {
message: result.failure_message().unwrap_or_default(),
name_test: result.name_test,
})
}
async fn run<C: TestTransaction>(
env_file: &str,
trace: &TraceSink,
handler: impl AsyncFnOnce(&C) -> Result<(), C::Error>,
) -> (Reason<C::Error>, Vec<Issue>) {
let (_connection, db, timeout) = match open::<C>(env_file, trace).await {
Ok(opened) => opened,
Err(setup) => return (Reason::Setup(setup), Vec::new()),
};
let outcome =
tokio::time::timeout(timeout, AssertUnwindSafe(handler(&db)).catch_unwind()).await;
let reason = match outcome {
Ok(Ok(Ok(()))) => Reason::Win,
Ok(Ok(Err(e))) => Reason::Error(e),
Ok(Err(payload)) => Reason::Panic(panic_message(payload.as_ref())),
Err(_) => Reason::TimedOut(timeout.as_secs()),
};
let mut issues = Vec::new();
if matches!(reason, Reason::Win) {
let swallowed = lock(trace)
.iter()
.filter(|q| q.failed && !q.expected)
.count();
if swallowed > 0 {
issues.push(Issue::SwallowedFailures(swallowed));
}
}
if !matches!(reason, Reason::TimedOut(_)) {
let before = lock(trace).len();
if let Ok(Ok(false)) = tokio::time::timeout(CLEANUP_TIMEOUT, db.still_open()).await {
issues.push(Issue::TransactionEnded);
}
lock(trace).truncate(before);
}
match tokio::time::timeout(CLEANUP_TIMEOUT, db.rollback_test()).await {
Ok(Ok(())) => {}
Ok(Err(e)) => issues.push(Issue::RollbackFailed(e.to_string())),
Err(_) => issues.push(Issue::RollbackFailed(msgf(
"runique_test.rollback_timed_out",
&[CLEANUP_TIMEOUT.as_secs()],
))),
}
(reason, issues)
}
async fn open<C: TestTransaction>(
env_file: &str,
trace: &TraceSink,
) -> Result<(C::Connection, C, Duration), String> {
let path = resolve(env_file);
let vars = read_env_file(&path)?;
let timeout = timeout_from(&vars)?;
let config = C::load_config(&vars).map_err(|e| e.to_string())?;
announce(&format!("{} ({})", C::describe(&config), path.display()));
let connection = C::connect(&config, trace.clone())
.await
.map_err(|e| msgf("runique_test.connect_failed", &[e]))?;
let db = C::begin_test(&connection)
.await
.map_err(|e| msgf("runique_test.begin_failed", &[e]))?;
Ok((connection, db, timeout))
}
fn announce(target: &str) {
let mut last = LAST_TARGET.lock().unwrap_or_else(|e| e.into_inner());
if last.as_deref() != Some(target) {
println!("\nrunique test → {target}");
*last = Some(target.to_string());
}
}
fn resolve(env_file: &str) -> PathBuf {
let path = Path::new(env_file);
match std::env::var_os("CARGO_MANIFEST_DIR") {
Some(dir) if path.is_relative() => Path::new(&dir).join(path),
_ => path.to_path_buf(),
}
}
fn timeout_from(vars: &StrMap) -> Result<Duration, String> {
match vars.get(TIMEOUT_KEY) {
None => Ok(DEFAULT_TIMEOUT),
Some(raw) => match raw.trim().parse::<u64>() {
Ok(secs) if secs > 0 => Ok(Duration::from_secs(secs)),
_ => Err(msgf(
"runique_test.bad_timeout",
&[TIMEOUT_KEY, raw.as_str()],
)),
},
}
}
pub(crate) fn record(sql: &str, elapsed: Duration, failed: bool) {
let _ = CURRENT.try_with(|scope| {
lock(&scope.trace).push(QueryTrace {
sql: sql.to_string(),
elapsed,
failed,
expected: false,
});
});
}
pub(crate) fn trace_len() -> usize {
CURRENT
.try_with(|scope| lock(&scope.trace).len())
.unwrap_or(0)
}
pub(crate) fn mark_expected_from(start: usize) {
let _ = CURRENT.try_with(|scope| {
for query in lock(&scope.trace).iter_mut().skip(start) {
query.expected = true;
}
});
}
fn lock(trace: &TraceSink) -> std::sync::MutexGuard<'_, Vec<QueryTrace>> {
trace.lock().unwrap_or_else(|e| e.into_inner())
}
fn panic_message(payload: &(dyn Any + Send)) -> String {
payload
.downcast_ref::<&str>()
.map(|s| s.to_string())
.or_else(|| payload.downcast_ref::<String>().cloned())
.unwrap_or_else(|| msg("runique_test.no_panic_message").into_owned())
}
fn current_test_name() -> String {
let name = std::thread::current()
.name()
.unwrap_or("unnamed test")
.to_string();
name.strip_prefix("runique_test::")
.map(str::to_string)
.unwrap_or(name)
}
fn read_env_file(path: &Path) -> Result<StrMap, String> {
let shown = path.display().to_string();
dotenvy::from_filename_iter(path)
.map_err(|e| {
msgf(
"runique_test.env_read_failed",
&[shown.as_str(), e.to_string().as_str()],
)
})?
.map(|item| item.map_err(|e| parse_error(&shown, &e)))
.collect()
}
fn parse_error(path: &str, error: &dotenvy::Error) -> String {
match error {
dotenvy::Error::LineParse(_, index) => msgf(
"runique_test.env_line_malformed",
&[path, index.to_string().as_str()],
),
other => msgf(
"runique_test.env_parse_failed",
&[path, other.to_string().as_str()],
),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn absolute_path_is_left_alone() {
let absolute = std::env::temp_dir().join("x.env");
assert_eq!(resolve(absolute.to_str().unwrap()), absolute);
}
#[test]
fn timeout_defaults_and_parses() {
let mut vars = StrMap::new();
assert_eq!(timeout_from(&vars).unwrap(), DEFAULT_TIMEOUT);
vars.insert(TIMEOUT_KEY.into(), "3".into());
assert_eq!(timeout_from(&vars).unwrap(), Duration::from_secs(3));
for bad in ["0", "-1", "soon", ""] {
vars.insert(TIMEOUT_KEY.into(), bad.into());
assert!(timeout_from(&vars).is_err(), "{bad:?} should be refused");
}
}
#[test]
fn a_relative_env_file_is_read_from_the_package() {
let dir = std::env::var_os("CARGO_MANIFEST_DIR").expect("set by cargo");
assert_eq!(resolve(".env.test"), Path::new(&dir).join(".env.test"));
}
#[test]
fn panic_messages_and_test_names() {
assert_eq!(panic_message(&"static text"), "static text");
assert_eq!(panic_message(&String::from("owned text")), "owned text");
assert_eq!(panic_message(&42u8), msg("runique_test.no_panic_message"));
assert_eq!(
current_test_name(),
"logic::builder_test::tests::panic_messages_and_test_names"
);
}
#[test]
fn a_malformed_env_line_is_never_quoted() {
let err = dotenvy::Error::LineParse("DATABASE_URL=postgres://u:secret@h/db".into(), 3);
let shown = parse_error("app/.env.test", &err);
assert!(
shown.contains("app/.env.test") && shown.contains('3'),
"{shown}"
);
assert!(!shown.contains("secret"), "{shown}");
let io = dotenvy::Error::Io(std::io::Error::other("disk gone"));
assert!(parse_error("app/.env.test", &io).contains("disk gone"));
}
#[test]
fn the_trace_is_empty_outside_a_test_and_counts_inside() {
assert_eq!(trace_len(), 0);
let scope = Arc::new(TestScope {
trace: Arc::default(),
});
let rt = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
let n = rt.block_on(CURRENT.scope(scope, async {
record("SELECT 1", Duration::ZERO, false);
record("SELECT 2", Duration::ZERO, false);
trace_len()
}));
assert_eq!(n, 2);
}
#[test]
fn a_new_target_is_remembered() {
announce("target-one");
assert_eq!(LAST_TARGET.lock().unwrap().as_deref(), Some("target-one"));
announce("target-two");
assert_eq!(LAST_TARGET.lock().unwrap().as_deref(), Some("target-two"));
}
}