use std::num::NonZeroU16;
use std::time::Duration;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum CacheEstimateModel {
Sol,
Terra,
Luna,
Astra,
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub(crate) struct CacheEstimateContext {
key: Option<NonZeroU16>,
}
impl CacheEstimateContext {
pub(crate) fn private(model: CacheEstimateModel, model_params: tau_proto::ModelParams) -> Self {
use tau_proto::{EffectiveReasoningEffort, ServiceTier};
let model = match model {
CacheEstimateModel::Sol => 1,
CacheEstimateModel::Terra => 2,
CacheEstimateModel::Luna => 3,
CacheEstimateModel::Astra => 4,
};
let (effective_kind, effective_level) = match model_params.effort.effective {
EffectiveReasoningEffort::ProviderDefault(None) => (0, 0),
EffectiveReasoningEffort::ProviderDefault(Some(level)) => (1, level as u16),
EffectiveReasoningEffort::Native(level) => (2, level as u16),
EffectiveReasoningEffort::Fixed(level) => (3, level as u16),
EffectiveReasoningEffort::Unsupported => (4, 0),
};
let service_tier = match model_params.service_tier {
None => 0,
Some(ServiceTier::Fast) => 1,
Some(ServiceTier::Flex) => 2,
};
let packed = model
| effective_kind << 3
| effective_level << 6
| u16::from(model_params.verbosity.as_u8()) << 9
| u16::from(model_params.thinking_summary.as_u8()) << 11
| service_tier << 13;
Self {
key: NonZeroU16::new(packed),
}
}
pub(crate) fn continues(self, previous: Self) -> bool {
self.key.is_some() && self == previous
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum CacheEstimateGeometry {
Step128Lag182,
Step1024Residue256,
Step1024Residue512,
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub(crate) enum CacheEstimateCalibration {
#[default]
Uncalibrated,
Candidate(CacheEstimateGeometry),
Confirmed(CacheEstimateGeometry),
}
impl CacheEstimateGeometry {
const MIN_OBSERVED_PREFIX_TOKENS: u64 = 10_000;
pub(crate) fn estimate(self, predecessor_input: u64) -> Option<u64> {
if predecessor_input < Self::MIN_OBSERVED_PREFIX_TOKENS {
return None;
}
match self {
Self::Step128Lag182 => Some(
predecessor_input
.saturating_sub(182)
.div_euclid(128)
.saturating_mul(128),
),
Self::Step1024Residue256 => Some(
predecessor_input
.saturating_sub(438)
.div_euclid(1_024)
.saturating_mul(1_024)
.saturating_add(256),
),
Self::Step1024Residue512 => Some(
predecessor_input
.saturating_sub(591)
.div_euclid(1_024)
.saturating_mul(1_024)
.saturating_add(512),
),
}
}
pub(crate) fn infer(predecessor_input: u64, cached_input: u64) -> Option<Self> {
if cached_input == 0 {
return None;
}
let mut matched = [
Self::Step128Lag182,
Self::Step1024Residue256,
Self::Step1024Residue512,
]
.into_iter()
.filter(|geometry| geometry.estimate(predecessor_input) == Some(cached_input));
let geometry = matched.next()?;
matched.next().is_none().then_some(geometry)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum CacheReadCeilingProjection {
Exact {
ceiling: u64,
context: CacheEstimateContext,
},
Estimated(CacheEstimateContext),
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct TurnStatsUsageProjection {
pub(crate) prompt_sent_tokens: u64,
pub(crate) prompt_cached_tokens: u64,
pub(crate) cache_read_ceiling: CacheReadCeilingProjection,
pub(crate) response_received_tokens: u64,
}
impl From<&tau_proto::ProviderTokenUsage> for TurnStatsUsageProjection {
fn from(usage: &tau_proto::ProviderTokenUsage) -> Self {
Self {
prompt_sent_tokens: usage.prompt_sent_tokens,
prompt_cached_tokens: usage.prompt_cached_tokens,
cache_read_ceiling: usage.prompt_cache_read_ceiling_tokens.map_or(
CacheReadCeilingProjection::Estimated(CacheEstimateContext::default()),
|ceiling| CacheReadCeilingProjection::Exact {
ceiling,
context: CacheEstimateContext::default(),
},
),
response_received_tokens: usage.response_received_tokens,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct PreviousTurnUsageProjection {
pub(crate) prompt_sent_tokens: u64,
pub(crate) response_received_tokens: u64,
pub(crate) cache_estimate_context: CacheEstimateContext,
pub(crate) cache_estimate_calibration: CacheEstimateCalibration,
}
impl From<(TurnStatsUsageProjection, CacheEstimateContext)> for PreviousTurnUsageProjection {
fn from(
(usage, cache_estimate_context): (TurnStatsUsageProjection, CacheEstimateContext),
) -> Self {
Self {
prompt_sent_tokens: usage.prompt_sent_tokens,
response_received_tokens: usage.response_received_tokens,
cache_estimate_context,
cache_estimate_calibration: CacheEstimateCalibration::Uncalibrated,
}
}
}
impl PreviousTurnUsageProjection {
pub(crate) fn from_completed(usage: TurnStatsUsageProjection, previous: Option<Self>) -> Self {
let cache_estimate_context = usage.estimate_context();
let cache_estimate_calibration = previous
.filter(|previous| cache_estimate_context.continues(previous.cache_estimate_context))
.map_or(CacheEstimateCalibration::Uncalibrated, |previous| {
let observed = CacheEstimateGeometry::infer(
previous.prompt_sent_tokens,
usage.prompt_cached_tokens,
);
match (previous.cache_estimate_calibration, observed) {
(
CacheEstimateCalibration::Candidate(expected)
| CacheEstimateCalibration::Confirmed(expected),
Some(observed),
) if expected == observed => CacheEstimateCalibration::Confirmed(observed),
(_, Some(observed)) => CacheEstimateCalibration::Candidate(observed),
_ => CacheEstimateCalibration::Uncalibrated,
}
});
Self {
prompt_sent_tokens: usage.prompt_sent_tokens,
response_received_tokens: usage.response_received_tokens,
cache_estimate_context,
cache_estimate_calibration,
}
}
}
impl TurnStatsUsageProjection {
pub(crate) fn with_estimate_context(mut self, context: CacheEstimateContext) -> Self {
self.cache_read_ceiling = match self.cache_read_ceiling {
CacheReadCeilingProjection::Exact { ceiling, .. } => {
CacheReadCeilingProjection::Exact { ceiling, context }
}
CacheReadCeilingProjection::Estimated(_) => {
CacheReadCeilingProjection::Estimated(context)
}
};
self
}
pub(crate) fn estimate_context(self) -> CacheEstimateContext {
match self.cache_read_ceiling {
CacheReadCeilingProjection::Exact { context, .. }
| CacheReadCeilingProjection::Estimated(context) => context,
}
}
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub(crate) struct CumulativeTurnUsageProjection {
pub(crate) sent_tokens: u64,
pub(crate) cached_tokens: u64,
pub(crate) received_tokens: u64,
}
impl From<tau_proto::TokenUsageCounts> for CumulativeTurnUsageProjection {
fn from(usage: tau_proto::TokenUsageCounts) -> Self {
Self {
sent_tokens: usage.sent_tokens,
cached_tokens: usage.cached_tokens,
received_tokens: usage.received_tokens,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct TurnStatsPresentationProjection {
pub(crate) finished_local_time: Option<(u8, u8)>,
pub(crate) usage: TurnStatsUsageProjection,
pub(crate) cumulative_usage: CumulativeTurnUsageProjection,
pub(crate) previous_usage: Option<PreviousTurnUsageProjection>,
pub(crate) turn_latency: Option<Duration>,
pub(crate) total_latency: Option<Duration>,
}