Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
237 changes: 84 additions & 153 deletions rust/src/providers/grok/billing/mod.rs
Original file line number Diff line number Diff line change
@@ -1,19 +1,29 @@
use chrono::{DateTime, TimeZone, Utc};
use reqwest::header::HeaderMap;

use super::product_usage::{GrokProductUsage, compose};
use crate::core::ProviderError;

mod protobuf;
mod reset_coupons;
mod web_product_usage;
#[cfg(test)]
pub(super) mod web_product_usage_tests;

use protobuf::{FieldValue, ProtobufField, looks_like_protobuf_payload};
pub(super) use reset_coupons::parse_grpc_web_reset_coupons;
use web_product_usage::decode_product_usage;

#[derive(Debug, Clone, Copy)]
#[derive(Debug, Clone)]
pub(super) struct GrokBillingSnapshot {
pub(super) used_percent: Option<f64>,
pub(super) used_percent_is_wire_published: bool,
pub(super) used_percent_is_implicit_zero: bool,
pub(super) resets_at: Option<DateTime<Utc>>,
pub(super) window_minutes: Option<u32>,
/// Product shares composing this same answer's wire percent; empty for
/// every other percent source.
pub(super) product_usage: Vec<GrokProductUsage>,
}

pub(super) fn validate_grpc_headers(headers: &HeaderMap) -> Result<(), ProviderError> {
Expand Down Expand Up @@ -46,26 +56,6 @@ pub(super) fn map_grpc_status(status: u16, context: &str) -> Result<(), Provider
Ok(())
}

/// Decode a length-prefixed field body: `read_varint -> try_from ->
/// checked_add -> bounds-check` in one place.
pub(super) fn read_length_field(
data: &[u8],
index: usize,
what: &str,
) -> Result<(usize, usize), ProviderError> {
let (len, start) = read_varint(data, index)
.ok_or_else(|| ProviderError::Parse(format!("Grok {what} is malformed")))?;
let len = usize::try_from(len)
.map_err(|_| ProviderError::Parse(format!("Grok {what} is too large")))?;
let end = start
.checked_add(len)
.ok_or_else(|| ProviderError::Parse(format!("Grok {what} length overflowed")))?;
if end > data.len() {
return Err(ProviderError::Parse(format!("Grok {what} is truncated")));
}
Ok((start, end))
}

/// Decode a varint Unix-seconds timestamp with the shared epoch bounding.
pub(super) fn unix_seconds_timestamp(seconds: u64) -> Option<DateTime<Utc>> {
// Varint timestamps are Unix seconds inside the range checked below.
Expand All @@ -81,8 +71,7 @@ pub(super) fn unix_seconds_timestamp(seconds: u64) -> Option<DateTime<Utc>> {
}

/// One parameterized gRPC-web frame walker. `on_malformed` decides the
/// malformed-frame policy: billing swallows malformed frames, the optional
/// reset lookup fails closed.
/// malformed-frame policy; both callers fail closed, with their own message.
///
/// Yields `(flags, payload)` for every frame, data and trailer alike; callers
/// split on the trailer flag.
Expand Down Expand Up @@ -114,20 +103,6 @@ pub(super) fn grpc_web_frames(
Ok(frames)
}

fn skip_field(data: &[u8], index: usize, wire: u64) -> Option<usize> {
match wire {
0 => read_varint(data, index).map(|(_, next)| next),
1 => index.checked_add(8).filter(|end| *end <= data.len()),
2 => {
let (len, start) = read_varint(data, index)?;
let len = usize::try_from(len).ok()?;
start.checked_add(len).filter(|end| *end <= data.len())
}
5 => index.checked_add(4).filter(|end| *end <= data.len()),
_ => None,
}
}

fn parse_grpc_web_response_at(
data: &[u8],
now: DateTime<Utc>,
Expand Down Expand Up @@ -167,11 +142,11 @@ fn parse_grpc_web_response_at(
.iter()
.any(|field| field.value != first.value)
});
let parsed_percent = if scan.is_complete && !conflicting_percent {
valid_percent_fields.first().map(|field| field.value as f64)
} else {
None
};
let percent_field = valid_percent_fields
.first()
.filter(|_| scan.is_complete && !conflicting_percent);
let parsed_percent = percent_field.map(|field| f64::from(field.value));
let percent_path = percent_field.map(|field| field.path.as_slice());

let reset_fields: Vec<(&VarintField, DateTime<Utc>)> = scan
.varints
Expand Down Expand Up @@ -203,24 +178,26 @@ fn parse_grpc_web_response_at(
let used_percent_is_implicit_zero =
no_usage_yet && payloads.len() == 1 && scan.is_complete && has_active_current_period;
let used_percent = parsed_percent.or_else(|| used_percent_is_implicit_zero.then_some(0.0));
// Shares are trusted only from one complete message whose aggregate is the
// config's own `[1, 1]` percent, never a nested or repeated lookalike.
let product_usage = match (parsed_percent, payloads.as_slice()) {
(Some(percent), [payload])
if scan.is_complete && percent_path == Some([1, 1].as_slice()) =>
{
compose(decode_product_usage(payload), percent)
}
_ => Vec::new(),
};
Ok(GrokBillingSnapshot {
used_percent,
used_percent_is_wire_published: parsed_percent.is_some(),
used_percent_is_implicit_zero,
resets_at,
window_minutes,
product_usage,
})
}

fn looks_like_protobuf_payload(data: &[u8]) -> bool {
let Some(&first) = data.first() else {
return false;
};
let field_number = first >> 3;
let wire_type = first & 0x07;
field_number > 0 && matches!(wire_type, 0 | 1 | 2 | 5)
}

fn varint_timestamp(field: &VarintField) -> Option<DateTime<Utc>> {
unix_seconds_timestamp(field.value)
}
Expand Down Expand Up @@ -258,9 +235,21 @@ fn unique_varint_at_path(scan: &ProtoScan, path: &[u64]) -> Option<u64> {
values.all(|value| value == first).then_some(first)
}

/// Data-frame payloads of a billing response. Empty when any frame is
/// truncated or carries a compressed or reserved flag: such a response must
/// never supply a partial percent, product shares, or an implicit zero.
fn grpc_web_data_frames(data: &[u8]) -> Vec<Vec<u8>> {
grpc_web_frames(data, |_| None)
.unwrap_or_default()
let Ok(frames) = grpc_web_frames(data, |context| {
Some(ProviderError::Parse(format!(
"Grok web billing frame is {context}"
)))
}) else {
return Vec::new();
};
if frames.iter().any(|(flags, _)| !matches!(flags, 0 | 0x80)) {
return Vec::new();
}
frames
.into_iter()
.filter(|(flags, _)| flags & 0x80 == 0)
.map(|(_, payload)| payload.to_vec())
Expand Down Expand Up @@ -297,120 +286,62 @@ struct VarintField {
}

impl ProtoScan {
/// Walk one message. The first malformed field ends the scan and marks it
/// incomplete, so nothing after it can be mistaken for a real field.
fn scan_message(&mut self, data: &[u8], path: &mut Vec<u64>, depth: usize) {
if depth > 8 {
self.is_complete = false;
return;
}
let mut i = 0;
while i < data.len() {
let field_start = i;
let Some((field, wire, next)) = read_key(data, i) else {
let mut index = 0;
while index < data.len() {
let Some(field) = ProtobufField::read(data, &mut index) else {
self.is_complete = false;
i = field_start.saturating_add(1);
continue;
return;
};
i = next;
path.push(field);
let Some(next) = self.scan_field(data, i, path, depth, wire) else {
self.is_complete = false;
path.pop();
i = field_start.saturating_add(1);
continue;
};
i = next;
path.push(field.number);
let accepted = self.scan_field(field.value, path, depth);
path.pop();
if !accepted {
self.is_complete = false;
return;
}
}
}

fn scan_field(
&mut self,
data: &[u8],
i: usize,
path: &mut Vec<u64>,
depth: usize,
wire: u64,
) -> Option<usize> {
if (path.as_slice() == [1, 1] && wire != 5) || (is_known_billing_message(path) && wire != 2)
{
return None;
}
match wire {
0 => self.scan_varint(data, i, path),
2 => self.scan_length_delimited(data, i, path, depth),
5 => self.scan_fixed32(data, i, path),
1 => i.checked_add(8).filter(|end| *end <= data.len()),
_ => None,
}
}

fn scan_varint(&mut self, data: &[u8], i: usize, path: &[u64]) -> Option<usize> {
let (value, next) = read_varint(data, i)?;
self.varints.push(VarintField {
path: path.to_vec(),
value,
});
Some(next)
}

fn scan_length_delimited(
&mut self,
data: &[u8],
i: usize,
path: &mut Vec<u64>,
depth: usize,
) -> Option<usize> {
let (len, next) = read_varint(data, i)?;
let start = next;
let len_usize = usize::try_from(len).ok()?;
let end = start.checked_add(len_usize)?;
if end > data.len() {
return None;
}
if depth < 4 && is_known_billing_message(path) {
self.scan_message(&data[start..end], path, depth + 1);
}
Some(end)
}

fn scan_fixed32(&mut self, data: &[u8], i: usize, path: &[u64]) -> Option<usize> {
let end = i.checked_add(4)?;
if end > data.len() {
return None;
}
let bytes = [data[i], data[i + 1], data[i + 2], data[i + 3]];
self.fixed32.push(Fixed32Field {
path: path.to_vec(),
value: f32::from_le_bytes(bytes),
order: self.order,
});
self.order += 1;
Some(end)
}
}

fn read_key(data: &[u8], i: usize) -> Option<(u64, u64, usize)> {
let (key, next) = read_varint(data, i)?;
let field = key >> 3;
(field > 0 && field <= 536_870_911).then_some((field, key & 0x07, next))
}

fn read_varint(data: &[u8], mut i: usize) -> Option<(u64, usize)> {
let mut value = 0u64;
let mut shift = 0;
while i < data.len() && shift < 64 {
let b = data[i];
i += 1;
if shift == 63 && b > 1 {
return None;
/// Record one field; `false` when its wire type contradicts the descriptor.
fn scan_field(&mut self, value: FieldValue<'_>, path: &mut Vec<u64>, depth: usize) -> bool {
let known_message = is_known_billing_message(path);
let wire_matches_descriptor = match value {
FieldValue::Message(_) => path.as_slice() != [1, 1],
FieldValue::Fixed32(_) => !known_message,
FieldValue::Varint(_) | FieldValue::Fixed64 => {
!known_message && path.as_slice() != [1, 1]
}
};
if !wire_matches_descriptor {
return false;
}
value |= u64::from(b & 0x7f) << shift;
if b & 0x80 == 0 {
return Some((value, i));
match value {
FieldValue::Varint(value) => self.varints.push(VarintField {
path: path.clone(),
value,
}),
FieldValue::Fixed32(value) => {
self.fixed32.push(Fixed32Field {
path: path.clone(),
value,
order: self.order,
});
self.order += 1;
}
FieldValue::Message(bytes) if depth < 4 && known_message => {
self.scan_message(bytes, path, depth + 1);
}
FieldValue::Message(_) | FieldValue::Fixed64 => {}
}
shift += 7;
true
}
None
}

fn is_known_billing_message(path: &[u64]) -> bool {
Expand Down
Loading