use crate::error::Result;
use serde_json::Value;
use std::{future::Future, pin::Pin};
pub struct ProviderRequest {
pub url: String,
pub headers: Vec<(String, String)>,
pub body: Value,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RequestAdmissionPolicy {
native_namespace: Option<String>,
protected_native_fields: &'static [&'static str],
native_prompt_fields: &'static [&'static str],
native_output_limit_fields: &'static [&'static str],
}
impl RequestAdmissionPolicy {
pub fn namespaced(
native_namespace: impl Into<String>,
protected_native_fields: &'static [&'static str],
native_prompt_fields: &'static [&'static str],
native_output_limit_fields: &'static [&'static str],
) -> Self {
Self {
native_namespace: Some(native_namespace.into()),
protected_native_fields,
native_prompt_fields,
native_output_limit_fields,
}
}
pub fn native_namespace(&self) -> Option<&str> {
self.native_namespace.as_deref()
}
pub fn protected_native_fields(&self) -> &'static [&'static str] {
self.protected_native_fields
}
pub fn native_prompt_fields(&self) -> &'static [&'static str] {
self.native_prompt_fields
}
pub fn native_output_limit_fields(&self) -> &'static [&'static str] {
self.native_output_limit_fields
}
}
impl ProviderRequest {
pub fn can_continue_from(&self, previous: &Self) -> bool {
self.url == previous.url
&& self.headers == previous.headers
&& crate::cache::continuation_matches(&previous.body, &self.body)
}
}
pub trait Provider: Send + Sync {
fn name(&self) -> &str;
fn request_admission_policy(&self) -> RequestAdmissionPolicy {
RequestAdmissionPolicy::default()
}
fn replay_target(&self, model: &str) -> crate::reasoning::ReplayTarget {
crate::reasoning::ReplayTarget::new(
self.name(),
model,
crate::reasoning::WireFormat::OpenAiChat,
)
}
fn request_replay_target(
&self,
model: &str,
request: &ProviderRequest,
) -> crate::reasoning::ReplayTarget {
self.replay_target(request.body["model"].as_str().unwrap_or(model))
}
fn transform_request(&self, model: &str, request: &Value) -> Result<ProviderRequest>;
fn prepare_request<'a>(
&'a self,
model: &'a str,
request: &'a Value,
) -> Pin<Box<dyn Future<Output = Result<ProviderRequest>> + Send + 'a>> {
Box::pin(async move { self.transform_request(model, request) })
}
fn transform_response(&self, model: &str, response: Value) -> Result<Value>;
fn stream_normalizer(&self, model: &str) -> crate::streaming::StreamNormalizer {
crate::streaming::StreamNormalizer::new(self.replay_target(model))
}
fn transform_stream_chunk(&self, model: &str, chunk: &str) -> Result<Option<String>>;
}