use serde::{Deserialize, Serialize};
use super::error::GraphError;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct TokenUsage {
pub input_tokens: u64,
pub output_tokens: u64,
}
impl TokenUsage {
#[must_use]
pub fn from_counts(input_tokens: Option<u64>, output_tokens: Option<u64>) -> Option<Self> {
match (input_tokens, output_tokens) {
(None, None) => None,
(input, output) => Some(Self {
input_tokens: input.unwrap_or(0),
output_tokens: output.unwrap_or(0),
}),
}
}
#[must_use]
pub fn total(self) -> u64 {
self.input_tokens + self.output_tokens
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Completion {
pub text: String,
pub usage: Option<TokenUsage>,
}
impl Completion {
#[must_use]
pub fn unmeasured(text: impl Into<String>) -> Self {
Self {
text: text.into(),
usage: None,
}
}
#[must_use]
pub fn measured(text: impl Into<String>, usage: Option<TokenUsage>) -> Self {
Self {
text: text.into(),
usage,
}
}
}
#[async_trait::async_trait]
pub trait LlmProvider: Send + Sync {
async fn complete(
&self,
system_prompt: &str,
user_message: &str,
max_tokens: u32,
) -> Result<String, GraphError>;
async fn complete_measured(
&self,
system_prompt: &str,
user_message: &str,
max_tokens: u32,
) -> Result<Completion, GraphError> {
let text = self
.complete(system_prompt, user_message, max_tokens)
.await?;
Ok(Completion::unmeasured(text))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_provider_that_reports_nothing_measures_nothing() {
assert_eq!(TokenUsage::from_counts(None, None), None);
}
#[test]
fn one_reported_side_is_still_a_measurement() {
assert_eq!(
TokenUsage::from_counts(None, Some(5)),
Some(TokenUsage {
input_tokens: 0,
output_tokens: 5,
})
);
}
#[test]
fn total_sums_both_sides() {
let usage = TokenUsage::from_counts(Some(13_658), Some(5)).unwrap();
assert_eq!(usage.total(), 13_663);
}
}