use crate::{
error::Result,
messages::Message,
model::ModelConfig,
providers::{
ProviderExt,
gemini::GeminiProvider,
types::{ChatResponse, ProviderSource},
},
};
#[derive(Debug)]
pub struct LLM {
pub provider_source: ProviderSource,
pub provider: Box<dyn ProviderExt>,
pub config: ModelConfig,
}
impl LLM {
pub fn new(provider_source: ProviderSource, model_name: String) -> Self {
let default_model_config = ModelConfig::new(&model_name);
let provider: Box<dyn ProviderExt> = match provider_source {
ProviderSource::Gemini => Box::new(GeminiProvider::with_default_config()),
_ => panic!(
"Unsupported provider source: {:?}. Supported providers: Gemini",
provider_source
),
};
LLM {
provider_source,
provider,
config: default_model_config,
}
}
pub fn gemini<S: Into<String>>(model_name: S) -> Self {
Self::new(ProviderSource::Gemini, model_name.into())
}
pub fn conservative(provider_source: ProviderSource, model_name: String) -> Self {
let config = ModelConfig::conservative(&model_name);
Self::new(provider_source, model_name).with_custom_config(config)
}
pub fn creative(provider_source: ProviderSource, model_name: String) -> Self {
let config = ModelConfig::creative(&model_name);
Self::new(provider_source, model_name).with_custom_config(config)
}
pub fn balanced(provider_source: ProviderSource, model_name: String) -> Self {
let config = ModelConfig::balanced(&model_name);
Self::new(provider_source, model_name).with_custom_config(config)
}
pub fn with_custom_config(mut self, config: ModelConfig) -> Self {
self.config = config;
self
}
pub fn temperature(&mut self, temperature: f32) -> &mut Self {
self.config.temperature = temperature;
self
}
pub fn system_instruction(&mut self, system_instruction: String) -> &mut Self {
self.config.system_instruction = Some(system_instruction);
self
}
pub fn get_config(&self) -> &ModelConfig {
&self.config
}
pub fn get_provider_source(&self) -> &ProviderSource {
&self.provider_source
}
pub fn get_model_name(&self) -> &str {
&self.config.name
}
pub async fn prompt<S: Into<String>>(&self, prompt: S) -> Result<ChatResponse> {
let config = self.config.clone();
self.provider.prompt(config, prompt.into()).await
}
pub async fn chat(&self, message: Message, history: Vec<Message>) -> Result<ChatResponse> {
let config = self.config.clone();
self.provider.chat(config, message, history).await
}
pub fn provider_name(&self) -> &'static str {
self.provider.name()
}
pub fn supports_streaming(&self) -> bool {
self.provider.supports_streaming()
}
pub fn supports_tools(&self) -> bool {
self.provider.supports_tools()
}
}
#[cfg(test)]
mod tests {
use crate::providers::gemini;
use super::*;
#[tokio::test]
async fn test_llm_creation() {
let _llm = LLM::new(
ProviderSource::Gemini,
gemini::PREDEFINED_MODELS[0].to_string(),
);
let _conservative = LLM::conservative(
ProviderSource::Gemini,
gemini::PREDEFINED_MODELS[0].to_string(),
);
let _creative = LLM::creative(
ProviderSource::Gemini,
gemini::PREDEFINED_MODELS[0].to_string(),
);
let _balanced = LLM::balanced(
ProviderSource::Gemini,
gemini::PREDEFINED_MODELS[0].to_string(),
);
}
}