use std::sync::Arc;
use anyhow::Result;
use async_trait::async_trait;
use crate::llm_client::LlmClient;
use crate::llm_client::StreamEventBox;
use crate::models::{MessageRequest, MessageResponse};
#[async_trait]
#[allow(dead_code)]
pub trait ModelClient: Send + Sync {
fn provider_name(&self) -> &str;
fn model(&self) -> &str;
fn billing_base_url(&self) -> Option<&str> {
None
}
fn effective_route_envelope(
&self,
requested_model: &str,
dispatched_at: chrono::DateTime<chrono::Utc>,
) -> crate::cost_status::EffectiveRouteEnvelope {
let provider = crate::config::ApiProvider::parse(self.provider_name())
.unwrap_or(crate::config::ApiProvider::Custom);
crate::cost_status::EffectiveRouteEnvelope::capture(
None,
provider,
self.provider_name(),
requested_model,
self.billing_base_url(),
dispatched_at,
)
}
async fn create_message(&self, request: MessageRequest) -> Result<MessageResponse>;
async fn create_message_stream(&self, request: MessageRequest) -> Result<StreamEventBox>;
async fn health_check(&self) -> Result<bool>;
}
pub type SharedModelClient = Arc<dyn ModelClient>;
#[async_trait]
impl<T> ModelClient for T
where
T: LlmClient + Send + Sync,
{
fn provider_name(&self) -> &str {
LlmClient::provider_name(self)
}
fn model(&self) -> &str {
LlmClient::model(self)
}
fn billing_base_url(&self) -> Option<&str> {
LlmClient::billing_base_url(self)
}
fn effective_route_envelope(
&self,
requested_model: &str,
dispatched_at: chrono::DateTime<chrono::Utc>,
) -> crate::cost_status::EffectiveRouteEnvelope {
LlmClient::effective_route_envelope(self, requested_model, dispatched_at)
}
async fn create_message(&self, request: MessageRequest) -> Result<MessageResponse> {
LlmClient::create_message(self, request).await
}
async fn create_message_stream(&self, request: MessageRequest) -> Result<StreamEventBox> {
LlmClient::create_message_stream(self, request).await
}
async fn health_check(&self) -> Result<bool> {
LlmClient::health_check(self).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn model_client_is_object_safe() {
fn accepts_dyn(_: Option<SharedModelClient>) {}
accepts_dyn(None);
}
}