pub mod obfuscate;
pub use obfuscate::{FnObfuscator, MapObfuscator, ObfuscatingCompletionModel, Obfuscator};
use crate::error::CompletionError;
use crate::message::Message;
use crate::telemetry::{classify_completion_error, CompletionTiming, LlmMetrics};
use crate::tool::ToolDefinition;
use serde::{de::DeserializeOwned, Deserialize, Serialize};
use std::future::Future;
use std::time::Instant;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct Usage {
pub input_tokens: u64,
pub output_tokens: u64,
#[serde(default)]
pub cache_read_tokens: u64,
#[serde(default)]
pub cache_creation_tokens: u64,
}
impl Usage {
pub fn new(input_tokens: u64, output_tokens: u64) -> Self {
Self {
input_tokens,
output_tokens,
cache_read_tokens: 0,
cache_creation_tokens: 0,
}
}
pub fn with_cache(
input_tokens: u64,
output_tokens: u64,
cache_read_tokens: u64,
cache_creation_tokens: u64,
) -> Self {
Self {
input_tokens,
output_tokens,
cache_read_tokens,
cache_creation_tokens,
}
}
pub fn total(&self) -> u64 {
self.input_tokens + self.output_tokens
}
pub fn cache_total(&self) -> u64 {
self.cache_read_tokens + self.cache_creation_tokens
}
}
impl std::ops::Add for Usage {
type Output = Self;
fn add(self, rhs: Self) -> Self {
Self {
input_tokens: self.input_tokens + rhs.input_tokens,
output_tokens: self.output_tokens + rhs.output_tokens,
cache_read_tokens: self.cache_read_tokens + rhs.cache_read_tokens,
cache_creation_tokens: self.cache_creation_tokens + rhs.cache_creation_tokens,
}
}
}
impl std::ops::AddAssign for Usage {
fn add_assign(&mut self, rhs: Self) {
self.input_tokens += rhs.input_tokens;
self.output_tokens += rhs.output_tokens;
self.cache_read_tokens += rhs.cache_read_tokens;
self.cache_creation_tokens += rhs.cache_creation_tokens;
}
}
#[derive(Debug, Clone, Default)]
pub struct CompletionRequest {
pub preamble: Option<String>,
pub messages: Vec<Message>,
pub tools: Vec<ToolDefinition>,
pub temperature: Option<f64>,
pub max_tokens: Option<u64>,
pub additional_params: Option<serde_json::Value>,
}
impl CompletionRequest {
pub fn new(message: impl Into<Message>) -> Self {
Self {
messages: vec![message.into()],
..Default::default()
}
}
pub fn with_preamble(mut self, preamble: impl Into<String>) -> Self {
self.preamble = Some(preamble.into());
self
}
pub fn with_messages(mut self, messages: Vec<Message>) -> Self {
self.messages = messages;
self
}
pub fn add_message(mut self, message: impl Into<Message>) -> Self {
self.messages.push(message.into());
self
}
pub fn with_tools(mut self, tools: Vec<ToolDefinition>) -> Self {
self.tools = tools;
self
}
pub fn with_temperature(mut self, temperature: f64) -> Self {
self.temperature = Some(temperature);
self
}
pub fn with_max_tokens(mut self, max_tokens: u64) -> Self {
self.max_tokens = Some(max_tokens);
self
}
pub fn with_additional_params(mut self, params: serde_json::Value) -> Self {
self.additional_params = Some(params);
self
}
}
#[derive(Debug, Clone)]
pub struct CompletionResponse<R = serde_json::Value> {
pub message: Message,
pub usage: Usage,
pub raw: R,
pub reasoning_content: Option<String>,
pub finish_reason: Option<String>,
}
impl<R> CompletionResponse<R> {
pub fn new(message: Message, usage: Usage, raw: R) -> Self {
Self {
message,
usage,
raw,
reasoning_content: None,
finish_reason: None,
}
}
pub fn with_reasoning(
message: Message,
usage: Usage,
raw: R,
reasoning: Option<String>,
) -> Self {
Self {
message,
usage,
raw,
reasoning_content: reasoning,
finish_reason: None,
}
}
pub fn is_truncated(&self) -> bool {
matches!(
self.finish_reason.as_deref(),
Some("length") | Some("max_tokens")
)
}
pub fn content(&self) -> String {
self.message.text()
}
pub fn has_tool_calls(&self) -> bool {
self.message.has_tool_calls()
}
pub fn tool_calls(&self) -> Vec<&crate::message::ToolCall> {
self.message.tool_calls()
}
}
pub trait CompletionModel: Clone + Send + Sync + 'static {
type Response: Send + Sync + Serialize + DeserializeOwned + 'static;
fn completion(
&self,
request: CompletionRequest,
) -> impl Future<Output = Result<CompletionResponse<Self::Response>, CompletionError>> + Send;
fn model_id(&self) -> &str;
fn provider(&self) -> &str;
}
pub struct CompletionRequestBuilder<M: CompletionModel> {
model: M,
request: CompletionRequest,
}
impl<M: CompletionModel> CompletionRequestBuilder<M> {
pub fn new(model: M, prompt: impl Into<Message>) -> Self {
Self {
model,
request: CompletionRequest::new(prompt),
}
}
pub fn preamble(mut self, preamble: impl Into<String>) -> Self {
self.request.preamble = Some(preamble.into());
self
}
pub fn messages(mut self, messages: Vec<Message>) -> Self {
self.request.messages = messages;
self
}
pub fn add_message(mut self, message: impl Into<Message>) -> Self {
self.request.messages.push(message.into());
self
}
pub fn tools(mut self, tools: Vec<ToolDefinition>) -> Self {
self.request.tools = tools;
self
}
pub fn temperature(mut self, temp: f64) -> Self {
self.request.temperature = Some(temp);
self
}
pub fn max_tokens(mut self, max: u64) -> Self {
self.request.max_tokens = Some(max);
self
}
pub fn additional_params(mut self, params: serde_json::Value) -> Self {
self.request.additional_params = Some(params);
self
}
pub async fn send(self) -> Result<CompletionResponse<M::Response>, CompletionError> {
self.model.completion(self.request).await
}
}
#[derive(Clone)]
pub struct MetricsCompletionModel<M: CompletionModel> {
inner: M,
metrics: LlmMetrics,
}
impl<M: CompletionModel> MetricsCompletionModel<M> {
pub fn new(inner: M) -> Self {
Self {
inner,
metrics: LlmMetrics::global(),
}
}
pub fn with_custom_metrics(inner: M, metrics: LlmMetrics) -> Self {
Self { inner, metrics }
}
pub fn inner(&self) -> &M {
&self.inner
}
pub fn into_inner(self) -> M {
self.inner
}
}
impl<M: CompletionModel> CompletionModel for MetricsCompletionModel<M> {
type Response = M::Response;
async fn completion(
&self,
request: CompletionRequest,
) -> Result<CompletionResponse<Self::Response>, CompletionError> {
let start = Instant::now();
let result = self.inner.completion(request).await;
let latency = start.elapsed();
let timing = match &result {
Ok(response) => CompletionTiming {
latency_ms: latency.as_secs_f64() * 1000.0,
input_tokens: response.usage.input_tokens,
output_tokens: response.usage.output_tokens,
provider: self.inner.provider().to_string(),
model: self.inner.model_id().to_string(),
success: true,
error_type: None,
},
Err(e) => CompletionTiming {
latency_ms: latency.as_secs_f64() * 1000.0,
input_tokens: 0,
output_tokens: 0,
provider: self.inner.provider().to_string(),
model: self.inner.model_id().to_string(),
success: false,
error_type: Some(classify_completion_error(e)),
},
};
self.metrics.record(&timing);
result
}
fn model_id(&self) -> &str {
self.inner.model_id()
}
fn provider(&self) -> &str {
self.inner.provider()
}
}
pub trait CompletionModelExt: CompletionModel + Sized {
fn with_metrics(self) -> MetricsCompletionModel<Self> {
MetricsCompletionModel::new(self)
}
fn with_obfuscator<O: Obfuscator>(self, obfuscator: O) -> ObfuscatingCompletionModel<Self, O> {
ObfuscatingCompletionModel::new(self, obfuscator)
}
}
impl<M: CompletionModel> CompletionModelExt for M {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_usage_operations() {
let u1 = Usage::new(100, 50);
let u2 = Usage::new(200, 100);
assert_eq!(u1.total(), 150);
assert_eq!((u1 + u2).total(), 450);
let mut u3 = Usage::new(10, 5);
u3 += Usage::new(20, 10);
assert_eq!(u3.total(), 45);
}
#[test]
fn test_completion_request_builder() {
let request = CompletionRequest::new("Hello")
.with_preamble("You are helpful")
.with_temperature(0.7)
.with_max_tokens(100);
assert_eq!(request.preamble, Some("You are helpful".to_string()));
assert_eq!(request.temperature, Some(0.7));
assert_eq!(request.max_tokens, Some(100));
assert_eq!(request.messages.len(), 1);
}
}