#[cfg(feature = "providers")]
pub mod anthropic;
#[cfg(feature = "providers")]
mod anthropic_stream;
#[cfg(feature = "providers")]
pub mod openai;
#[cfg(feature = "providers")]
mod openai_stream;
#[cfg(feature = "providers")]
mod sse;
#[cfg(feature = "providers")]
mod wire;
use std::fmt::Debug;
use std::sync::Arc;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::core::{
Disposition, Effect, EffectDescriptor, EffectError, Recovery, RetryPolicy, Sensitivity, Spend,
Trust,
};
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub struct ModelId {
pub provider: String,
pub model: String,
}
impl ModelId {
pub fn new(provider: impl Into<String>, model: impl Into<String>) -> Self {
Self {
provider: provider.into(),
model: model.into(),
}
}
}
impl std::fmt::Display for ModelId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}/{}", self.provider, self.model)
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct Usage {
pub input_tokens: u64,
pub output_tokens: u64,
#[serde(default)]
pub cache_write_tokens: u64,
#[serde(default)]
pub cache_read_tokens: u64,
pub minor_units: i64,
}
impl Usage {
#[must_use]
pub const fn spend(&self) -> Spend {
Spend {
tokens: self.input_tokens + self.output_tokens,
minor_units: self.minor_units,
}
}
#[must_use]
pub const fn uncached_input_tokens(&self) -> u64 {
self.input_tokens
.saturating_sub(self.cache_write_tokens)
.saturating_sub(self.cache_read_tokens)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Completion {
pub text: String,
pub usage: Usage,
pub stop_reason: Option<String>,
pub truncated: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub structured: Option<Value>,
}
#[derive(Debug, thiserror::Error)]
pub enum ModelError {
#[error("could not reach '{model}': {detail}")]
Unreachable { model: ModelId, detail: String },
#[error("'{model}' refused the request: {detail}")]
Refused { model: ModelId, detail: String },
#[error("'{model}' is rate limiting: {detail}")]
RateLimited { model: ModelId, detail: String },
#[error("'{model}' stopped mid-response after {} token(s): {detail}", usage.input_tokens + usage.output_tokens)]
Interrupted {
model: ModelId,
usage: Usage,
detail: String,
},
#[error("'{model}' did not say whether it generated: {detail}")]
Unavailable { model: ModelId, detail: String },
#[error("'{model}' generated and then died without saying what it cost: {detail}")]
Unaccounted { model: ModelId, detail: String },
#[error("'{model}' returned an unusable answer: {detail}")]
Unusable {
model: ModelId,
usage: Usage,
detail: String,
},
}
impl ModelError {
#[must_use]
pub const fn disposition(&self) -> Disposition {
match self {
Self::Unreachable { .. }
| Self::Refused { .. }
| Self::RateLimited { .. }
| Self::Unavailable { .. } => Disposition::DidNotHappen,
Self::Interrupted { .. } | Self::Unusable { .. } | Self::Unaccounted { .. } => {
Disposition::Landed
}
}
}
#[must_use]
pub const fn usage(&self) -> Usage {
match self {
Self::Interrupted { usage, .. } | Self::Unusable { usage, .. } => *usage,
_ => Usage {
input_tokens: 0,
output_tokens: 0,
cache_write_tokens: 0,
cache_read_tokens: 0,
minor_units: 0,
},
}
}
}
#[derive(Debug, Clone)]
pub struct Request<'a> {
pub model: &'a ModelId,
pub prompt: &'a Value,
pub schema: Option<&'a Value>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SchemaMode {
#[default]
Native,
ForcedTool,
}
#[async_trait]
pub trait ModelProvider: Send + Sync + Debug {
async fn complete(&self, request: Request<'_>) -> Result<Completion, ModelError>;
}
#[derive(Debug)]
pub struct ModelCall {
model: ModelId,
prompt: Value,
schema: Option<Value>,
provider: Arc<dyn ModelProvider>,
max_sensitivity: Sensitivity,
output_sensitivity: Sensitivity,
retry: RetryPolicy,
}
impl ModelCall {
#[must_use]
pub fn new(provider: Arc<dyn ModelProvider>, model: ModelId, prompt: Value) -> Self {
Self {
model,
prompt,
schema: None,
provider,
max_sensitivity: Sensitivity::Public,
output_sensitivity: Sensitivity::Public,
retry: RetryPolicy::never(),
}
}
#[must_use]
pub const fn with_max_sensitivity(mut self, s: Sensitivity) -> Self {
self.max_sensitivity = s;
self
}
#[must_use]
pub const fn with_output_sensitivity(mut self, s: Sensitivity) -> Self {
self.output_sensitivity = s;
self
}
#[must_use]
pub const fn with_retry(mut self, r: RetryPolicy) -> Self {
self.retry = r;
self
}
#[must_use]
pub fn expecting(mut self, schema: Value) -> Self {
self.schema = Some(schema);
self
}
}
#[async_trait]
impl Effect for ModelCall {
type Output = Completion;
fn descriptor(&self) -> EffectDescriptor {
EffectDescriptor::new(
"model.complete",
serde_json::json!({
"provider": self.model.provider,
"model": self.model.model,
"prompt": self.prompt,
"schema": self.schema,
}),
)
}
fn mutates(&self) -> bool {
false
}
fn recovery(&self) -> Recovery {
Recovery::Retry
}
fn retry(&self) -> RetryPolicy {
self.retry
}
fn max_sensitivity(&self) -> Sensitivity {
self.max_sensitivity
}
fn output_sensitivity(&self) -> Sensitivity {
self.output_sensitivity
}
fn trust(&self) -> Trust {
Trust::Untrusted
}
fn spend(&self, output: &Completion) -> Spend {
output.usage.spend()
}
async fn perform(&self) -> Result<Completion, EffectError> {
self.provider
.complete(Request {
model: &self.model,
prompt: &self.prompt,
schema: self.schema.as_ref(),
})
.await
.map_err(|e| {
let detail = e.to_string();
let spend = e.usage().spend();
if spend.is_zero() {
match e.disposition() {
Disposition::DidNotHappen => EffectError::Rejected(detail),
Disposition::InDoubt => EffectError::Interrupted {
driver: self.model.to_string(),
detail,
},
Disposition::Landed => EffectError::Performed(detail),
}
} else {
EffectError::Metered {
detail,
spend,
disposition: e.disposition(),
}
}
})
}
}