use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use arrow::array::{Array, BooleanArray, StringArray};
use arrow::record_batch::RecordBatch;
use datafusion::prelude::SessionContext;
use serde_json::Value;
use tracing::field::{Field, Visit};
use crate::sources::providers::open_connector::json_to_arrow::RowConverter;
use crate::sources::providers::open_connector::row_path::RowPath;
use crate::sources::providers::open_connector::source_pack::SourcePackTable;
pub(crate) use skardi_source_pack::testing::{
MockGateway, MockResponse, RecordedRequest, discovery_ok, envelope_err, envelope_ok,
};
pub(crate) async fn collect(ctx: &SessionContext, sql: &str) -> Vec<RecordBatch> {
ctx.sql(sql)
.await
.unwrap_or_else(|e| panic!("failed to plan {sql}: {e}"))
.collect()
.await
.unwrap_or_else(|e| panic!("failed to collect {sql}: {e}"))
}
pub(crate) fn convert_first_page(table: &SourcePackTable, page: &Value) -> RecordBatch {
let rows = RowPath::parse(table.row_path)
.expect("row path")
.rows(page, 1)
.expect("row array");
RowConverter::new(table.fields)
.expect("converter")
.convert(rows, 1)
.expect("page converts")
}
pub(crate) fn utf8<'a>(batch: &'a RecordBatch, name: &str) -> &'a StringArray {
let column = batch
.column_by_name(name)
.unwrap_or_else(|| panic!("column {name}"));
column
.as_any()
.downcast_ref()
.unwrap_or_else(|| panic!("column {name} is {:?}, not Utf8", column.data_type()))
}
pub(crate) fn boolean<'a>(batch: &'a RecordBatch, name: &str) -> &'a BooleanArray {
let column = batch
.column_by_name(name)
.unwrap_or_else(|| panic!("column {name}"));
column
.as_any()
.downcast_ref()
.unwrap_or_else(|| panic!("column {name} is {:?}, not Boolean", column.data_type()))
}
pub(crate) fn column_values(batches: &[RecordBatch], name: &str) -> Vec<String> {
let mut offset = 0;
batches
.iter()
.flat_map(|batch| {
let values = utf8(batch, name).clone();
let base = offset;
offset += values.len();
(0..values.len()).map(move |i| {
assert!(
!values.is_null(i),
"column {name} is NULL at row {}",
base + i
);
values.value(i).to_string()
})
})
.collect()
}
pub(crate) fn execute_inputs(gateway: &MockGateway, action_path: &str) -> Vec<Value> {
gateway
.requests()
.into_iter()
.filter(|r| r.method == "POST" && r.path.ends_with(action_path))
.map(|r| {
serde_json::from_str::<Value>(&r.body).expect("request body is JSON")["input"].clone()
})
.collect()
}
pub(crate) fn input_keys(input: &Value) -> Vec<&str> {
let mut keys: Vec<&str> = input
.as_object()
.expect("input object")
.keys()
.map(String::as_str)
.collect();
keys.sort_unstable();
keys
}
#[derive(Debug, Clone)]
pub(crate) struct CapturedEvent {
pub(crate) level: tracing::Level,
pub(crate) message: String,
pub(crate) fields: HashMap<String, String>,
}
impl CapturedEvent {
pub(crate) fn field(&self, name: &str) -> Option<&str> {
self.fields.get(name).map(String::as_str)
}
}
#[derive(Default)]
struct FieldRecorder {
message: String,
fields: HashMap<String, String>,
}
impl Visit for FieldRecorder {
fn record_str(&mut self, field: &Field, value: &str) {
self.fields
.insert(field.name().to_string(), value.to_string());
}
fn record_bool(&mut self, field: &Field, value: bool) {
self.fields
.insert(field.name().to_string(), value.to_string());
}
fn record_u64(&mut self, field: &Field, value: u64) {
self.fields
.insert(field.name().to_string(), value.to_string());
}
fn record_i64(&mut self, field: &Field, value: i64) {
self.fields
.insert(field.name().to_string(), value.to_string());
}
fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
let rendered = format!("{value:?}");
if field.name() == "message" {
self.message = rendered;
} else {
self.fields.insert(field.name().to_string(), rendered);
}
}
}
struct CaptureSubscriber {
events: Arc<Mutex<Vec<CapturedEvent>>>,
next_span_id: AtomicU64,
}
impl tracing::Subscriber for CaptureSubscriber {
fn enabled(&self, _metadata: &tracing::Metadata<'_>) -> bool {
true
}
fn new_span(&self, _attrs: &tracing::span::Attributes<'_>) -> tracing::span::Id {
tracing::span::Id::from_u64(self.next_span_id.fetch_add(1, Ordering::Relaxed))
}
fn record(&self, _span: &tracing::span::Id, _values: &tracing::span::Record<'_>) {}
fn record_follows_from(&self, _span: &tracing::span::Id, _follows: &tracing::span::Id) {}
fn event(&self, event: &tracing::Event<'_>) {
let mut recorder = FieldRecorder::default();
event.record(&mut recorder);
self.events
.lock()
.unwrap_or_else(|p| p.into_inner())
.push(CapturedEvent {
level: *event.metadata().level(),
message: recorder.message,
fields: recorder.fields,
});
}
fn enter(&self, _span: &tracing::span::Id) {}
fn exit(&self, _span: &tracing::span::Id) {}
}
pub(crate) fn fingerprint_uncovered_columns(
contract: &str,
row_path: &str,
fields: &[crate::sources::providers::open_connector::json_to_arrow::FieldMapping],
) -> Vec<&'static str> {
fn descend<'a>(mut node: &'a serde_json::Value, path: &str) -> &'a serde_json::Value {
for segment in path.split('.') {
node = &node["properties"][segment];
}
node
}
let contract: serde_json::Value = serde_json::from_str(contract).expect("contract parses");
let row_schema: &serde_json::Value = if row_path == "$" {
&contract
} else {
&descend(&contract, row_path.strip_prefix("$.").expect("row path"))["items"]
};
fields
.iter()
.filter(|field| descend(row_schema, field.path).is_null())
.map(|field| field.name)
.collect()
}
pub(crate) struct EnvVarGuard {
name: String,
previous: Option<std::ffi::OsString>,
}
impl EnvVarGuard {
pub(crate) fn set(name: &str, value: &str) -> Self {
let previous = std::env::var_os(name);
unsafe { std::env::set_var(name, value) };
Self {
name: name.to_string(),
previous,
}
}
}
impl Drop for EnvVarGuard {
fn drop(&mut self) {
match self.previous.take() {
Some(value) => unsafe { std::env::set_var(&self.name, value) },
None => unsafe { std::env::remove_var(&self.name) },
}
}
}
pub(crate) fn capture_events() -> (
tracing::subscriber::DefaultGuard,
Arc<Mutex<Vec<CapturedEvent>>>,
) {
let events = Arc::new(Mutex::new(Vec::new()));
let subscriber = CaptureSubscriber {
events: Arc::clone(&events),
next_span_id: AtomicU64::new(1),
};
(tracing::subscriber::set_default(subscriber), events)
}