Skip to content
Merged
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
40 changes: 38 additions & 2 deletions src/analyzers/piebald.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@

use crate::analyzer::{Analyzer, DataSource};
use crate::contribution_cache::ContributionStrategy;
use crate::models::calculate_total_cost;
use crate::models::{InputTokenSemantics, calculate_total_cost, get_model_info};
use crate::types::{Application, ConversationMessage, MessageRole, Stats};
use crate::utils::hash_text;
use anyhow::Result;
Expand Down Expand Up @@ -150,6 +150,17 @@ fn parse_timestamp(ts: &str) -> Option<DateTime<Utc>> {
.map(|dt| dt.with_timezone(&Utc))
}

fn normalize_input_tokens(model: Option<&str>, input_tokens: u64, cache_read_tokens: u64) -> u64 {
if model
.and_then(get_model_info)
.is_some_and(|info| info.input_token_semantics == InputTokenSemantics::IncludesCacheRead)
{
input_tokens.saturating_sub(cache_read_tokens)
} else {
input_tokens
}
}

/// Convert Piebald messages to splitrail's ConversationMessage format.
fn convert_messages(
chats: &[PiebaldChat],
Expand Down Expand Up @@ -196,11 +207,13 @@ fn convert_messages(
};

// Map token stats
let input_tokens = msg.input_tokens.unwrap_or(0) as u64;
let raw_input_tokens = msg.input_tokens.unwrap_or(0) as u64;
let output_tokens = msg.output_tokens.unwrap_or(0) as u64;
let reasoning_tokens = msg.reasoning_tokens.unwrap_or(0) as u64;
let cache_read_tokens = msg.cache_read_tokens.unwrap_or(0) as u64;
let cache_creation_tokens = msg.cache_write_tokens.unwrap_or(0) as u64;
let input_tokens =
normalize_input_tokens(model_str.as_deref(), raw_input_tokens, cache_read_tokens);

// Calculate cost using splitrail's model pricing
let cost = if let Some(ref model) = model_str {
Expand Down Expand Up @@ -332,6 +345,29 @@ mod tests {
assert!(result.unwrap().messages.is_empty());
}

#[test]
fn test_normalize_input_tokens_subtracts_openai_cached_reads() {
assert_eq!(normalize_input_tokens(Some("gpt-5"), 1_000, 300), 700);
}

#[test]
fn test_normalize_input_tokens_subtracts_tiered_openai_cached_reads() {
assert_eq!(normalize_input_tokens(Some("gpt-5.5"), 1_000, 300), 700);
}

#[test]
fn test_normalize_input_tokens_preserves_anthropic_input() {
assert_eq!(
normalize_input_tokens(Some("claude-sonnet-4-20250514"), 700, 300),
700
);
}

#[test]
fn test_normalize_input_tokens_saturates_for_openai() {
assert_eq!(normalize_input_tokens(Some("gpt-5"), 100, 300), 0);
}

#[test]
fn test_parse_timestamp_rfc3339() {
let ts = "2025-12-10T14:30:00Z";
Expand Down
81 changes: 57 additions & 24 deletions src/models.rs
Original file line number Diff line number Diff line change
Expand Up @@ -53,20 +53,30 @@ pub struct TieredCaching {
pub bracket_pricing: bool,
}

/// Different caching support models
/// Different cache pricing structures.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum CachingSupport {
/// Model does not support caching
/// Model does not support caching.
None,
/// OpenAI-style caching (simple cached input pricing)
/// Flat cached input pricing.
OpenAI { cached_input_per_1m: f64 },
/// Anthropic-style caching (separate write and read costs)
/// Separate cache write and cache read pricing.
Anthropic {
cache_write_per_1m: f64,
cache_read_per_1m: f64,
},
/// Google-style caching (may have tiers like input/output)
Google(TieredCaching),
/// Tiered cached input pricing.
Tiered(TieredCaching),
}

/// How a provider reports input tokens relative to cache reads.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum InputTokenSemantics {
/// Reported input tokens exclude cache reads and can be added to cached tokens.
#[default]
ExcludesCache,
/// Reported input tokens include cache reads; subtract cache reads before normalizing.
IncludesCacheRead,
}

/// Complete model information with all pricing details
Expand All @@ -76,6 +86,9 @@ pub struct ModelInfo {
pub pricing: PricingStructure,
/// Caching support and pricing
pub caching: CachingSupport,
/// How provider usage reports input tokens relative to cache reads.
#[serde(default)]
pub input_token_semantics: InputTokenSemantics,
/// Whether pricing is estimated (not officially published by provider)
pub is_estimated: bool,
}
Expand Down Expand Up @@ -122,7 +135,7 @@ impl Registry {
};

let caching_ok = match &info.caching {
CachingSupport::Google(tiered) => {
CachingSupport::Tiered(tiered) => {
Self::validate_tier_bounds(&tiered.tiers, |tier| tier.max_tokens)
}
_ => true,
Expand Down Expand Up @@ -172,6 +185,19 @@ fn get_registry_lock() -> &'static RwLock<Registry> {
REGISTRY.get_or_init(|| RwLock::new(Registry::new_with_defaults()))
}

fn input_token_semantics_for_model(model_name: &str) -> InputTokenSemantics {
let name = model_name
.rsplit_once('/')
.map(|(_, name)| name)
.unwrap_or(model_name);

if name.starts_with("gpt-") || name.starts_with('o') {
InputTokenSemantics::IncludesCacheRead
} else {
InputTokenSemantics::ExcludesCache
}
}

fn populate_defaults(
index: &mut HashMap<String, Arc<ModelInfo>>,
aliases: &mut HashMap<String, String>,
Expand All @@ -183,6 +209,7 @@ fn populate_defaults(
Arc::new(ModelInfo {
pricing: $pricing,
caching: $caching,
input_token_semantics: input_token_semantics_for_model($name),
is_estimated: $est,
}),
);
Expand Down Expand Up @@ -526,7 +553,7 @@ fn populate_defaults(
],
bracket_pricing: false,
}),
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![
CachingTier {
max_tokens: Some(272_000),
Expand Down Expand Up @@ -580,7 +607,7 @@ fn populate_defaults(
],
bracket_pricing: false,
}),
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![
CachingTier {
max_tokens: Some(272_000),
Expand Down Expand Up @@ -818,7 +845,7 @@ fn populate_defaults(
input_per_1m: 0.5,
output_per_1m: 3.0
},
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![CachingTier {
max_tokens: None,
cached_input_per_1m: 0.05
Expand All @@ -844,7 +871,7 @@ fn populate_defaults(
],
bracket_pricing: true,
}),
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![
CachingTier {
max_tokens: Some(200_000),
Expand Down Expand Up @@ -896,7 +923,7 @@ fn populate_defaults(
],
bracket_pricing: false,
}),
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![
CachingTier {
max_tokens: Some(200_000),
Expand All @@ -917,7 +944,7 @@ fn populate_defaults(
input_per_1m: 0.3,
output_per_1m: 2.5
},
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![CachingTier {
max_tokens: None,
cached_input_per_1m: 0.075
Expand All @@ -932,7 +959,7 @@ fn populate_defaults(
input_per_1m: 0.1,
output_per_1m: 0.4
},
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![CachingTier {
max_tokens: None,
cached_input_per_1m: 0.025
Expand All @@ -947,7 +974,7 @@ fn populate_defaults(
input_per_1m: 0.0,
output_per_1m: 0.0
},
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![CachingTier {
max_tokens: None,
cached_input_per_1m: 0.0
Expand All @@ -962,7 +989,7 @@ fn populate_defaults(
input_per_1m: 0.1,
output_per_1m: 0.4
},
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![CachingTier {
max_tokens: None,
cached_input_per_1m: 0.025
Expand Down Expand Up @@ -997,7 +1024,7 @@ fn populate_defaults(
],
bracket_pricing: false,
}),
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![
CachingTier {
max_tokens: Some(128_000),
Expand Down Expand Up @@ -1029,7 +1056,7 @@ fn populate_defaults(
],
bracket_pricing: false,
}),
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![
CachingTier {
max_tokens: Some(128_000),
Expand Down Expand Up @@ -1061,7 +1088,7 @@ fn populate_defaults(
],
bracket_pricing: false,
}),
CachingSupport::Google(TieredCaching {
CachingSupport::Tiered(TieredCaching {
tiers: vec![
CachingTier {
max_tokens: Some(128_000),
Expand Down Expand Up @@ -1538,6 +1565,7 @@ fn get_free_model_info() -> Arc<ModelInfo> {
output_per_1m: 0.0,
},
caching: CachingSupport::None,
input_token_semantics: InputTokenSemantics::ExcludesCache,
is_estimated: false,
})
}))
Expand Down Expand Up @@ -1677,7 +1705,7 @@ fn cache_cost_for_model(
let read_cost = (cache_read_tokens as f64 / 1_000_000.0) * cache_read_per_1m;
creation_cost + read_cost
}
CachingSupport::Google(tiered) => {
CachingSupport::Tiered(tiered) => {
calculate_tiered_cache_cost(cache_read_tokens, &tiered.tiers, tiered.bracket_pricing)
}
}
Expand Down Expand Up @@ -1821,8 +1849,8 @@ where
#[cfg(test)]
mod tests {
use super::{
CachingSupport, CachingTier, ModelInfo, PricingStructure, PricingTier, Registry,
TieredCaching, TieredPricing, calculate_cache_cost, calculate_input_cost,
CachingSupport, CachingTier, InputTokenSemantics, ModelInfo, PricingStructure, PricingTier,
Registry, TieredCaching, TieredPricing, calculate_cache_cost, calculate_input_cost,
calculate_output_cost, get_model_info, get_registry_lock, init_external_models,
};

Expand Down Expand Up @@ -1858,6 +1886,7 @@ mod tests {
output_per_1m: 2000.0,
},
caching: CachingSupport::None,
input_token_semantics: InputTokenSemantics::default(),
is_estimated: false,
},
);
Expand Down Expand Up @@ -1897,6 +1926,7 @@ mod tests {
output_per_1m: 2.0,
},
caching: CachingSupport::None,
input_token_semantics: InputTokenSemantics::default(),
is_estimated: false,
},
);
Expand All @@ -1918,6 +1948,7 @@ mod tests {
output_per_1m: 4.0,
},
caching: CachingSupport::None,
input_token_semantics: InputTokenSemantics::default(),
is_estimated: false,
},
);
Expand Down Expand Up @@ -1948,6 +1979,7 @@ mod tests {
output_per_1m: 2.5,
},
caching: CachingSupport::None,
input_token_semantics: InputTokenSemantics::default(),
is_estimated: false,
},
);
Expand Down Expand Up @@ -1994,7 +2026,7 @@ mod tests {
],
bracket_pricing: false,
}),
caching: CachingSupport::Google(TieredCaching {
caching: CachingSupport::Tiered(TieredCaching {
tiers: vec![
CachingTier {
max_tokens: Some(50),
Expand All @@ -2007,6 +2039,7 @@ mod tests {
],
bracket_pricing: false,
}),
input_token_semantics: InputTokenSemantics::default(),
is_estimated: false,
},
);
Expand Down Expand Up @@ -2144,7 +2177,7 @@ mod tests {
}

match (&first.caching, &second.caching) {
(CachingSupport::Google(first_tiered), CachingSupport::Google(second_tiered)) => {
(CachingSupport::Tiered(first_tiered), CachingSupport::Tiered(second_tiered)) => {
assert!(
std::ptr::eq(first_tiered.tiers.as_ptr(), second_tiered.tiers.as_ptr()),
"cache tiers should not be reallocated on each lookup"
Expand Down
Loading