use af_context::{ProviderAttemptId, ToolCallId};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MeteringSource {
Reported,
Estimated,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MeteringOutcome {
Completed,
Failed,
Cancelled,
Steered,
TimedOut,
NotDispatched,
Unknown,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
pub enum MeteringDetails {
Model {
model: String,
provider: Option<String>,
provider_attempt_id: ProviderAttemptId,
source: MeteringSource,
outcome: MeteringOutcome,
},
Tool {
call_id: ToolCallId,
name: String,
source: MeteringSource,
outcome: MeteringOutcome,
},
}
impl MeteringDetails {
pub fn completes(&self, prepared: &Self) -> bool {
match (self, prepared) {
(
Self::Model {
model,
provider,
provider_attempt_id,
..
},
Self::Model {
model: expected_model,
provider: expected_provider,
provider_attempt_id: expected_attempt,
..
},
) => {
model == expected_model
&& provider_attempt_id == expected_attempt
&& expected_provider
.as_ref()
.is_none_or(|expected| provider.as_ref() == Some(expected))
}
(
Self::Tool { call_id, name, .. },
Self::Tool {
call_id: expected_call,
name: expected_name,
..
},
) => call_id == expected_call && name == expected_name,
_ => false,
}
}
pub fn validate(&self) -> Result<(), crate::EventError> {
let valid =
|s: &str| !s.trim().is_empty() && s.len() <= 512 && !s.chars().any(char::is_control);
let ok = match self {
Self::Model {
model,
provider,
provider_attempt_id,
..
} => {
valid(model)
&& provider.as_deref().is_none_or(valid)
&& valid(provider_attempt_id.as_str())
}
Self::Tool { call_id, name, .. } => valid(call_id.as_str()) && valid(name),
};
if ok {
Ok(())
} else {
Err(crate::EventError::InvalidMetering)
}
}
}
impl MeteringSource {
pub fn as_str(self) -> &'static str {
match self {
Self::Reported => "reported",
Self::Estimated => "estimated",
}
}
}
impl MeteringOutcome {
pub fn as_str(self) -> &'static str {
match self {
Self::Completed => "completed",
Self::Failed => "failed",
Self::Cancelled => "cancelled",
Self::Steered => "steered",
Self::TimedOut => "timed_out",
Self::NotDispatched => "not_dispatched",
Self::Unknown => "unknown",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn attribution_is_bounded_and_preserves_unknown_provider_and_outcomes() {
for outcome in [
MeteringOutcome::Completed,
MeteringOutcome::Failed,
MeteringOutcome::Cancelled,
MeteringOutcome::Steered,
MeteringOutcome::TimedOut,
MeteringOutcome::NotDispatched,
MeteringOutcome::Unknown,
] {
let mut details = MeteringDetails::Model {
model: "m".into(),
provider: None,
provider_attempt_id: "attempt".parse().unwrap(),
source: MeteringSource::Estimated,
outcome,
};
details.validate().unwrap();
assert_eq!(serde_json::to_value(outcome).unwrap(), outcome.as_str());
assert_eq!(
serde_json::from_value::<MeteringDetails>(serde_json::to_value(&details).unwrap())
.unwrap(),
details
);
if let MeteringDetails::Model { provider, .. } = &mut details {
*provider = Some("\n".into());
}
assert!(details.validate().is_err());
}
for name in ["".to_owned(), "a".repeat(513), "a\nb".into()] {
assert!(MeteringDetails::Tool {
name,
call_id: "call".parse().unwrap(),
source: MeteringSource::Estimated,
outcome: MeteringOutcome::Unknown
}
.validate()
.is_err());
}
}
}