use std::fmt;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
use crate::schema::ResponseShape;
#[derive(Debug, Clone, PartialEq)]
pub enum ModelSelection {
Exact(String),
Capabilities {
capabilities: Vec<String>,
min_context_tokens: Option<i64>,
},
Default,
}
#[derive(Debug, Clone)]
pub struct CompletionRequest {
pub node: String,
pub model: ModelSelection,
pub system: Option<String>,
pub prompt: String,
pub context: Vec<(String, Value)>,
pub response_type: String,
pub shape: ResponseShape,
pub max_tokens: u32,
}
impl CompletionRequest {
pub fn digest(&self) -> String {
let mut hasher = Sha256::new();
hasher.update(self.node.as_bytes());
hasher.update([0]);
hasher.update(self.system.as_deref().unwrap_or("").as_bytes());
hasher.update([0]);
hasher.update(self.prompt.as_bytes());
hasher.update([0]);
hasher.update(self.response_type.as_bytes());
for (name, value) in &self.context {
hasher.update([0]);
hasher.update(name.as_bytes());
hasher.update([0]);
hasher.update(value.to_string().as_bytes());
}
format!("{:x}", hasher.finalize())
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct CompletionResponse {
pub value: Value,
pub usage: Usage,
pub model: String,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct Usage {
pub input_tokens: u64,
pub output_tokens: u64,
#[serde(default, skip_serializing_if = "is_zero")]
pub cache_read_tokens: u64,
}
fn is_zero(value: &u64) -> bool {
*value == 0
}
impl Usage {
pub fn total(&self) -> u64 {
self.input_tokens + self.output_tokens
}
pub fn add(&mut self, other: Usage) {
self.input_tokens += other.input_tokens;
self.output_tokens += other.output_tokens;
self.cache_read_tokens += other.cache_read_tokens;
}
}
#[derive(Debug)]
pub enum ProviderError {
Transport(String),
Request { status: u16, message: String },
RateLimited { retry_after_seconds: Option<u64> },
Refused {
category: Option<String>,
explanation: Option<String>,
},
InvalidResponse(String),
Truncated { limit: u32 },
Configuration(String),
Cassette(String),
}
impl fmt::Display for ProviderError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ProviderError::Transport(message) => write!(f, "provider transport failed: {message}"),
ProviderError::Request { status, message } => {
write!(f, "provider rejected the request ({status}): {message}")
}
ProviderError::RateLimited {
retry_after_seconds: Some(seconds),
} => {
write!(f, "provider is rate limiting; retry after {seconds}s")
}
ProviderError::RateLimited { .. } => write!(f, "provider is rate limiting"),
ProviderError::Refused {
category,
explanation,
} => {
write!(f, "the provider declined to answer")?;
if let Some(category) = category {
write!(f, " ({category})")?;
}
if let Some(explanation) = explanation {
write!(f, ": {explanation}")?;
}
Ok(())
}
ProviderError::InvalidResponse(message) => {
write!(f, "the response did not match the declared type: {message}")
}
ProviderError::Truncated { limit } => {
write!(f, "the response was cut off at the {limit} token limit")
}
ProviderError::Configuration(message) => {
write!(f, "provider not configured: {message}")
}
ProviderError::Cassette(message) => write!(f, "cassette replay failed: {message}"),
}
}
}
impl std::error::Error for ProviderError {}
pub type DeltaSink<'a> = &'a mut dyn FnMut(&str);
pub trait ModelProvider {
fn name(&self) -> &str;
fn complete(
&mut self,
request: &CompletionRequest,
) -> Result<CompletionResponse, ProviderError>;
fn streams(&self) -> bool {
false
}
fn complete_streaming(
&mut self,
request: &CompletionRequest,
on_delta: DeltaSink<'_>,
) -> Result<CompletionResponse, ProviderError> {
let _ = on_delta;
self.complete(request)
}
}
impl<P: ModelProvider + ?Sized> ModelProvider for Box<P> {
fn name(&self) -> &str {
(**self).name()
}
fn complete(
&mut self,
request: &CompletionRequest,
) -> Result<CompletionResponse, ProviderError> {
(**self).complete(request)
}
fn streams(&self) -> bool {
(**self).streams()
}
fn complete_streaming(
&mut self,
request: &CompletionRequest,
on_delta: DeltaSink<'_>,
) -> Result<CompletionResponse, ProviderError> {
(**self).complete_streaming(request, on_delta)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn request() -> CompletionRequest {
CompletionRequest {
node: "n0".into(),
model: ModelSelection::Default,
system: None,
prompt: "Summarise this".into(),
context: vec![("document".into(), json!("hello"))],
response_type: "markdown".into(),
shape: ResponseShape::Prose,
max_tokens: 4096,
}
}
#[test]
fn the_digest_is_stable() {
assert_eq!(request().digest(), request().digest());
}
#[test]
fn the_digest_changes_with_the_prompt() {
let mut other = request();
other.prompt = "Summarise this document".into();
assert_ne!(request().digest(), other.digest());
}
#[test]
fn the_digest_changes_with_the_context() {
let mut other = request();
other.context = vec![("document".into(), json!("goodbye"))];
assert_ne!(request().digest(), other.digest());
}
#[test]
fn the_digest_changes_with_the_response_type() {
let mut other = request();
other.response_type = "text".into();
assert_ne!(request().digest(), other.digest());
}
#[test]
fn a_provider_that_cannot_stream_answers_at_once_and_shows_nothing() {
struct AtOnce;
impl ModelProvider for AtOnce {
fn name(&self) -> &str {
"at-once"
}
fn complete(
&mut self,
_request: &CompletionRequest,
) -> Result<CompletionResponse, ProviderError> {
Ok(CompletionResponse {
value: json!("the whole answer"),
usage: Usage::default(),
model: "at-once".into(),
})
}
}
let mut seen = Vec::new();
let response = AtOnce
.complete_streaming(&request(), &mut |text| seen.push(text.to_string()))
.unwrap();
assert_eq!(response.value, json!("the whole answer"));
assert!(
seen.is_empty(),
"a provider with nothing live to show must not invent deltas: {seen:?}"
);
assert!(!AtOnce.streams());
}
#[test]
fn usage_accumulates() {
let mut total = Usage::default();
total.add(Usage {
input_tokens: 10,
output_tokens: 5,
cache_read_tokens: 0,
});
total.add(Usage {
input_tokens: 1,
output_tokens: 2,
cache_read_tokens: 3,
});
assert_eq!(total.input_tokens, 11);
assert_eq!(total.output_tokens, 7);
assert_eq!(total.cache_read_tokens, 3);
assert_eq!(total.total(), 18);
}
}