1use super::provider::LLMError;
2use super::types::LLMResponse;
3use async_trait::async_trait;
4
5#[async_trait]
11pub trait LLMClient: Send + Sync {
12 async fn generate(&mut self, prompt: &str) -> Result<LLMResponse, LLMError>;
13 fn model_id(&self) -> &str;
14}
15
16pub type AnyClient = Box<dyn LLMClient>;
18
19pub struct ProviderClientAdapter {
23 provider: Box<dyn super::provider::LLMProvider>,
24 model_id: String,
25}
26
27impl ProviderClientAdapter {
28 pub fn new(provider: Box<dyn super::provider::LLMProvider>, model_id: String) -> Self {
30 Self { provider, model_id }
31 }
32}
33
34#[async_trait]
35impl LLMClient for ProviderClientAdapter {
36 async fn generate(&mut self, prompt: &str) -> Result<LLMResponse, LLMError> {
37 use super::provider::{LLMRequest, Message};
38 let request = LLMRequest {
39 messages: std::sync::Arc::new(vec![Message::user(prompt.to_string())]),
40 model: self.model_id.clone(),
41 ..Default::default()
42 };
43 Ok(self.provider.generate(request).await?)
44 }
45
46 fn model_id(&self) -> &str {
47 &self.model_id
48 }
49}