use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use axum::http::request::Parts;
use serde_json::Value;
const PROBE_TIMEOUT_SECS: u64 = 10;
#[derive(Clone, Debug, Default)]
pub struct CounterfactualSlot(Arc<AtomicU64>);
impl PartialEq for CounterfactualSlot {
fn eq(&self, other: &Self) -> bool {
self.get() == other.get()
}
}
impl CounterfactualSlot {
pub fn new() -> Self {
Self::default()
}
pub fn set(&self, tokens: u64) {
self.0.store(tokens.max(1), Ordering::Relaxed);
}
pub fn get(&self) -> Option<u64> {
match self.0.load(Ordering::Relaxed) {
0 => None,
n => Some(n),
}
}
}
const COUNT_TOKENS_FIELDS: &[&str] = &[
"model",
"messages",
"system",
"tools",
"tool_choice",
"thinking",
];
pub(crate) fn probe_body(original: &Value, original_model: Option<&str>) -> Option<Vec<u8>> {
let obj = original.as_object()?;
if !obj.contains_key("model") || !obj.contains_key("messages") {
return None;
}
let mut probe = serde_json::Map::new();
for &field in COUNT_TOKENS_FIELDS {
if let Some(v) = obj.get(field) {
probe.insert(field.to_string(), v.clone());
}
}
if let Some(model) = original_model {
probe.insert("model".to_string(), Value::String(model.to_string()));
}
serde_json::to_vec(&Value::Object(probe)).ok()
}
const PROBE_HEADERS: &[&str] = &[
"x-api-key",
"authorization",
"anthropic-version",
"anthropic-beta",
];
pub(crate) fn maybe_spawn_probe(
client: &reqwest::Client,
parts: &Parts,
upstream_base: &str,
original: Option<&Value>,
original_model: Option<&str>,
request_was_rewritten: bool,
) -> Option<CounterfactualSlot> {
if !request_was_rewritten
|| !parts
.uri
.path()
.trim_end_matches('/')
.ends_with("/v1/messages")
|| !crate::core::config::Config::load()
.proxy
.counterfactual_metering_enabled()
{
return None;
}
let body = probe_body(original?, original_model)?;
let url = format!(
"{}/v1/messages/count_tokens",
upstream_base.trim_end_matches('/')
);
let mut req = client
.post(&url)
.timeout(std::time::Duration::from_secs(PROBE_TIMEOUT_SECS))
.header("content-type", "application/json")
.body(body);
for &name in PROBE_HEADERS {
if let Some(v) = parts.headers.get(name) {
req = req.header(name, v.clone());
}
}
let slot = CounterfactualSlot::new();
let task_slot = slot.clone();
tokio::spawn(async move {
match req.send().await {
Ok(resp) if resp.status().is_success() => match resp.json::<Value>().await {
Ok(v) => {
if let Some(tokens) = v.get("input_tokens").and_then(Value::as_u64) {
task_slot.set(tokens);
} else {
tracing::debug!("counterfactual probe: response without input_tokens");
}
}
Err(e) => tracing::debug!("counterfactual probe: unreadable response: {e}"),
},
Ok(resp) => tracing::debug!(
"counterfactual probe: count_tokens returned {}",
resp.status()
),
Err(e) => tracing::debug!("counterfactual probe: {e}"),
}
});
Some(slot)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn full_request() -> Value {
json!({
"model": "claude-sonnet-4",
"messages": [{"role": "user", "content": "hello"}],
"system": "be terse",
"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}],
"tool_choice": {"type": "auto"},
"thinking": {"type": "enabled", "budget_tokens": 1024},
"max_tokens": 4096,
"stream": true,
"temperature": 0.7,
"metadata": {"user_id": "u1"}
})
}
#[test]
fn probe_body_is_the_count_tokens_whitelist() {
let body = probe_body(&full_request(), None).expect("probe body");
let v: Value = serde_json::from_slice(&body).unwrap();
let obj = v.as_object().unwrap();
for field in COUNT_TOKENS_FIELDS {
assert!(obj.contains_key(*field), "{field} must be projected");
}
for rejected in ["max_tokens", "stream", "temperature", "metadata"] {
assert!(!obj.contains_key(rejected), "{rejected} must be stripped");
}
assert_eq!(obj["model"], "claude-sonnet-4");
}
#[test]
fn probe_body_restores_the_prerouting_model() {
let mut req = full_request();
req["model"] = json!("claude-haiku-3.5");
let body = probe_body(&req, Some("claude-sonnet-4")).unwrap();
let v: Value = serde_json::from_slice(&body).unwrap();
assert_eq!(v["model"], "claude-sonnet-4");
}
#[test]
fn probe_body_requires_model_and_messages() {
assert!(probe_body(&json!({"messages": []}), None).is_none());
assert!(probe_body(&json!({"model": "m"}), None).is_none());
assert!(probe_body(&json!("not an object"), None).is_none());
}
#[test]
fn slot_roundtrip_and_zero_means_empty() {
let slot = CounterfactualSlot::new();
assert_eq!(slot.get(), None, "fresh slot is empty");
slot.set(1234);
assert_eq!(slot.get(), Some(1234));
let zero = CounterfactualSlot::new();
zero.set(0);
assert_eq!(zero.get(), Some(1));
}
#[test]
fn slots_share_state_across_clones() {
let slot = CounterfactualSlot::new();
let clone = slot.clone();
clone.set(77);
assert_eq!(slot.get(), Some(77));
}
}