use std::fmt;
use std::future::Future;
use std::pin::Pin;
use serde_json::Value;
pub trait Provider: Send + Sync {
fn capabilities(&self) -> Capabilities;
fn complete(
&self,
request: Request,
) -> impl Future<Output = Result<Response, ProviderError>> + Send;
}
type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
pub trait DynProvider: Send + Sync {
fn capabilities(&self) -> Capabilities;
fn complete_boxed(&self, request: Request) -> BoxFuture<'_, Result<Response, ProviderError>>;
}
impl<P: Provider> DynProvider for P {
fn capabilities(&self) -> Capabilities {
Provider::capabilities(self)
}
fn complete_boxed(&self, request: Request) -> BoxFuture<'_, Result<Response, ProviderError>> {
Box::pin(self.complete(request))
}
}
impl<P: Provider> Provider for &P {
fn capabilities(&self) -> Capabilities {
(**self).capabilities()
}
fn complete(
&self,
request: Request,
) -> impl Future<Output = Result<Response, ProviderError>> + Send {
(**self).complete(request)
}
}
impl<P: Provider> Provider for std::sync::Arc<P> {
fn capabilities(&self) -> Capabilities {
(**self).capabilities()
}
fn complete(
&self,
request: Request,
) -> impl Future<Output = Result<Response, ProviderError>> + Send {
(**self).complete(request)
}
}
impl Provider for Box<dyn DynProvider> {
fn capabilities(&self) -> Capabilities {
(**self).capabilities()
}
fn complete(
&self,
request: Request,
) -> impl Future<Output = Result<Response, ProviderError>> + Send {
(**self).complete_boxed(request)
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct Capabilities {
pub strategies: Vec<Strategy>,
pub dialect: SchemaDialect,
pub native_schema_visible: bool,
}
impl Capabilities {
pub fn new(strategies: Vec<Strategy>, dialect: SchemaDialect) -> Self {
Self {
strategies,
dialect,
native_schema_visible: true,
}
}
pub fn native_schema_visible(mut self, visible: bool) -> Self {
self.native_schema_visible = visible;
self
}
}
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Strategy {
NativeSchema,
ToolCall,
JsonMode,
PromptOnly,
}
impl Strategy {
pub fn as_str(self) -> &'static str {
match self {
Self::NativeSchema => "native_schema",
Self::ToolCall => "tool_call",
Self::JsonMode => "json_mode",
Self::PromptOnly => "prompt_only",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum SchemaDialect {
OpenAiStrict,
Generic,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct Request {
pub instructions: String,
pub input: Value,
pub output_schema: Value,
pub strategy: Strategy,
pub options: GenerationOptions,
#[serde(default)]
pub demonstrations: Vec<Demonstration>,
#[serde(default)]
pub tools: Vec<crate::tools::ToolSpec>,
pub repair: Vec<RepairTurn>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct Demonstration {
pub input: Value,
pub output: Value,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct RepairTurn {
pub response: String,
pub feedback: String,
}
#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct GenerationOptions {
pub temperature: Option<f32>,
pub max_tokens: Option<u32>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct Response {
pub content: String,
pub model: String,
pub finish: FinishReason,
pub usage: Usage,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum FinishReason {
Stop,
Length,
Refusal,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct Usage {
pub input_tokens: u64,
pub output_tokens: u64,
}
#[derive(Debug)]
#[non_exhaustive]
pub enum ProviderError {
Transport(String),
Status { code: u16, body: String },
InvalidResponse(String),
Recording(String),
}
impl fmt::Display for ProviderError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Transport(msg) => write!(f, "transport error: {msg}"),
Self::Status { code, body } => write!(f, "provider returned status {code}: {body}"),
Self::InvalidResponse(msg) => write!(f, "invalid provider response: {msg}"),
Self::Recording(msg) => write!(f, "recording: {msg}"),
}
}
}
impl std::error::Error for ProviderError {}