use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::BTreeMap;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU64, Ordering};
pub fn estimate_json_bytes(v: &Value) -> u64 {
match v {
Value::Null => 4,
Value::Bool(true) => 4,
Value::Bool(false) => 5,
Value::Number(n) => {
if let Some(i) = n.as_i64() {
digits_i64(i)
} else if let Some(u) = n.as_u64() {
digits_u64(u)
} else {
17
}
}
Value::String(s) => s.len() as u64 + 2,
Value::Array(items) => {
let inner: u64 = items.iter().map(estimate_json_bytes).sum();
inner + 2 + items.len().saturating_sub(1) as u64
}
Value::Object(map) => {
let inner: u64 = map
.iter()
.map(|(k, val)| k.len() as u64 + 3 + estimate_json_bytes(val))
.sum();
inner + 2 + map.len().saturating_sub(1) as u64
}
}
}
pub fn estimate_page_bytes(records: &[Value]) -> u64 {
records.iter().map(estimate_json_bytes).sum()
}
fn digits_u64(mut u: u64) -> u64 {
let mut n = 1;
while u >= 10 {
u /= 10;
n += 1;
}
n
}
fn digits_i64(i: i64) -> u64 {
if i < 0 {
1 + digits_u64(i.unsigned_abs())
} else {
digits_u64(i as u64)
}
}
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize, JsonSchema,
)]
#[serde(rename_all = "snake_case")]
pub enum UsageSide {
Source,
Sink,
}
impl UsageSide {
pub fn as_str(self) -> &'static str {
match self {
Self::Source => "source",
Self::Sink => "sink",
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
pub struct CostSignal {
pub kind: String,
pub unit: String,
pub quantity: f64,
pub side: UsageSide,
pub connector: String,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, JsonSchema)]
pub struct UsageSnapshot {
pub records_read: u64,
pub records_written: u64,
pub bytes_read: u64,
pub bytes_written: u64,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub source_roundtrips: BTreeMap<String, u64>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub sink_roundtrips: BTreeMap<String, u64>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub signals: Vec<CostSignal>,
#[serde(default, skip_serializing_if = "is_zero_u64")]
pub throttled: u64,
#[serde(default, skip_serializing_if = "is_zero_f64")]
pub throttle_wait_secs: f64,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub source_retries: BTreeMap<String, u64>,
}
fn is_zero_u64(v: &u64) -> bool {
*v == 0
}
fn is_zero_f64(v: &f64) -> bool {
*v == 0.0
}
impl UsageSnapshot {
pub fn roundtrips(&self, side: UsageSide) -> u64 {
match side {
UsageSide::Source => self.source_roundtrips.values().sum(),
UsageSide::Sink => self.sink_roundtrips.values().sum(),
}
}
pub fn merge(&mut self, other: &UsageSnapshot) {
self.records_read += other.records_read;
self.records_written += other.records_written;
self.bytes_read += other.bytes_read;
self.bytes_written += other.bytes_written;
for (k, v) in &other.source_roundtrips {
*self.source_roundtrips.entry(k.clone()).or_default() += v;
}
for (k, v) in &other.sink_roundtrips {
*self.sink_roundtrips.entry(k.clone()).or_default() += v;
}
self.signals.extend(other.signals.iter().cloned());
self.throttled += other.throttled;
self.throttle_wait_secs += other.throttle_wait_secs;
for (k, v) in &other.source_retries {
*self.source_retries.entry(k.clone()).or_default() += v;
}
}
}
#[derive(Debug, Default)]
pub struct UsageMeter {
records_read: AtomicU64,
records_written: AtomicU64,
bytes_read: AtomicU64,
bytes_written: AtomicU64,
roundtrips: Mutex<BTreeMap<(UsageSide, &'static str), u64>>,
signals: Mutex<Vec<CostSignal>>,
throttled: AtomicU64,
throttle_wait_nanos: AtomicU64,
source_retries: Mutex<BTreeMap<&'static str, u64>>,
}
impl UsageMeter {
pub fn new() -> Self {
Self::default()
}
pub fn add_read(&self, records: u64, bytes: u64) {
self.records_read.fetch_add(records, Ordering::Relaxed);
self.bytes_read.fetch_add(bytes, Ordering::Relaxed);
}
pub fn add_written(&self, records: u64, bytes: u64) {
self.records_written.fetch_add(records, Ordering::Relaxed);
self.bytes_written.fetch_add(bytes, Ordering::Relaxed);
}
pub fn add_roundtrip(&self, side: UsageSide, op: &'static str) {
let mut map = self.roundtrips.lock().unwrap_or_else(|e| e.into_inner());
*map.entry((side, op)).or_default() += 1;
}
pub fn add_signal(&self, signal: CostSignal) {
self.signals
.lock()
.unwrap_or_else(|e| e.into_inner())
.push(signal);
}
pub fn add_throttled(&self) {
self.throttled.fetch_add(1, Ordering::Relaxed);
}
pub fn add_throttle_wait(&self, slept: std::time::Duration) {
self.throttle_wait_nanos.fetch_add(
u64::try_from(slept.as_nanos()).unwrap_or(u64::MAX),
Ordering::Relaxed,
);
}
pub fn add_source_retry(&self, class: &'static str) {
let mut map = self
.source_retries
.lock()
.unwrap_or_else(|e| e.into_inner());
*map.entry(class).or_default() += 1;
}
pub fn records_written(&self) -> u64 {
self.records_written.load(Ordering::Relaxed)
}
pub fn bytes_written(&self) -> u64 {
self.bytes_written.load(Ordering::Relaxed)
}
pub fn snapshot(&self) -> UsageSnapshot {
let mut source_roundtrips = BTreeMap::new();
let mut sink_roundtrips = BTreeMap::new();
for ((side, op), n) in self
.roundtrips
.lock()
.unwrap_or_else(|e| e.into_inner())
.iter()
{
match side {
UsageSide::Source => *source_roundtrips.entry((*op).to_string()).or_default() += n,
UsageSide::Sink => *sink_roundtrips.entry((*op).to_string()).or_default() += n,
}
}
UsageSnapshot {
records_read: self.records_read.load(Ordering::Relaxed),
records_written: self.records_written.load(Ordering::Relaxed),
bytes_read: self.bytes_read.load(Ordering::Relaxed),
bytes_written: self.bytes_written.load(Ordering::Relaxed),
source_roundtrips,
sink_roundtrips,
signals: self
.signals
.lock()
.unwrap_or_else(|e| e.into_inner())
.clone(),
throttled: self.throttled.load(Ordering::Relaxed),
throttle_wait_secs: std::time::Duration::from_nanos(
self.throttle_wait_nanos.load(Ordering::Relaxed),
)
.as_secs_f64(),
source_retries: self
.source_retries
.lock()
.unwrap_or_else(|e| e.into_inner())
.iter()
.map(|(k, v)| ((*k).to_string(), *v))
.collect(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn byte_estimate_tracks_serialized_size() {
for v in [
json!(null),
json!(true),
json!(false),
json!(0),
json!(-42),
json!(1234567890123u64),
json!("héllo"),
json!([]),
json!({}),
json!([1, 2, 3]),
json!({"a": 1, "bb": [true, null], "c": {"d": "x"}}),
] {
let exact = serde_json::to_vec(&v).unwrap().len() as u64;
assert_eq!(estimate_json_bytes(&v), exact, "{v}");
}
assert!(estimate_json_bytes(&json!(1.5)) >= 3);
assert_eq!(
estimate_page_bytes(&[json!({"a": 1}), json!({"a": 22})]),
7 + 8
);
}
#[test]
fn meter_counts_and_snapshots() {
let m = UsageMeter::new();
m.add_read(3, 30);
m.add_written(2, 20);
m.add_roundtrip(UsageSide::Source, "page");
m.add_roundtrip(UsageSide::Source, "page");
m.add_roundtrip(UsageSide::Sink, "insert");
m.add_signal(CostSignal {
kind: "bytes_billed".into(),
unit: "bytes".into(),
quantity: 1024.0,
side: UsageSide::Sink,
connector: "bigquery".into(),
});
assert_eq!(m.records_written(), 2);
assert_eq!(m.bytes_written(), 20);
let s = m.snapshot();
assert_eq!(s.records_read, 3);
assert_eq!(s.bytes_read, 30);
assert_eq!(s.source_roundtrips["page"], 2);
assert_eq!(s.sink_roundtrips["insert"], 1);
assert_eq!(s.roundtrips(UsageSide::Source), 2);
assert_eq!(s.roundtrips(UsageSide::Sink), 1);
assert_eq!(s.signals.len(), 1);
assert_eq!(UsageSide::Sink.as_str(), "sink");
let mut total = UsageSnapshot::default();
total.merge(&s);
total.merge(&s);
assert_eq!(total.records_written, 4);
assert_eq!(total.source_roundtrips["page"], 4);
assert_eq!(total.signals.len(), 2);
let round: UsageSnapshot =
serde_json::from_value(serde_json::to_value(&s).unwrap()).unwrap();
assert_eq!(round, s);
}
#[test]
fn meter_tallies_throttling_and_retries() {
let m = UsageMeter::new();
let quiet = m.snapshot();
let v = serde_json::to_value(&quiet).unwrap();
assert!(v.get("throttled").is_none());
assert!(v.get("throttle_wait_secs").is_none());
assert!(v.get("source_retries").is_none());
m.add_throttled();
m.add_throttled();
m.add_throttle_wait(std::time::Duration::from_millis(1500));
m.add_throttle_wait(std::time::Duration::from_millis(500));
m.add_source_retry("rate_limited");
m.add_source_retry("rate_limited");
m.add_source_retry("http_5xx");
let s = m.snapshot();
assert_eq!(s.throttled, 2);
assert!((s.throttle_wait_secs - 2.0).abs() < 1e-9);
assert_eq!(s.source_retries["rate_limited"], 2);
assert_eq!(s.source_retries["http_5xx"], 1);
let mut total = UsageSnapshot::default();
total.merge(&s);
total.merge(&s);
assert_eq!(total.throttled, 4);
assert!((total.throttle_wait_secs - 4.0).abs() < 1e-9);
assert_eq!(total.source_retries["rate_limited"], 4);
let round: UsageSnapshot =
serde_json::from_value(serde_json::to_value(&s).unwrap()).unwrap();
assert_eq!(round, s);
}
}