Skip to main content

rpc_agent/providers/
mod.rs

1use std::fmt::Display;
2
3use rig::tool::Tool;
4use schemars::JsonSchema;
5
6use crate::{ToolWrapper, error::Error};
7
8mod ollama;
9mod openai;
10
11#[async_trait::async_trait]
12pub trait CompletionProvider: Send + Sync {
13    async fn chat(&self, prompt: &str) -> Result<String, Error>;
14    fn model(&self) -> &str;
15    fn api_key(&self) -> Option<&str>;
16    fn system_message(&self) -> Option<&str>;
17    fn temperature(&self) -> Option<f64>;
18    fn max_tokens(&self) -> Option<u64>;
19    fn provider(&self) -> Providers;
20}
21
22/// Supported AI providers for the agent server.
23pub enum Providers {
24    Ollama,
25    OpenAI,
26}
27
28impl Display for Providers {
29    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
30        match self {
31            Providers::Ollama => write!(f, "ollama"),
32            Providers::OpenAI => write!(f, "openai"),
33        }
34    }
35}
36
37impl From<&str> for Providers {
38    fn from(value: &str) -> Self {
39        match value.to_lowercase().as_str() {
40            "ollama" => Providers::Ollama,
41            "openai" => Providers::OpenAI,
42            _ => panic!(
43                "unknown provider: {}. Currently supported providers are: ollama, openai",
44                value
45            ),
46        }
47    }
48}
49
50impl Providers {
51    pub fn init<T: Tool + 'static>(
52        provider: Providers,
53        model: &str,
54        api_key: Option<&str>,
55        system_message: String,
56        temperature: Option<f64>,
57        max_tokens: Option<u64>,
58        tool: Option<ToolWrapper<T>>,
59    ) -> Result<Box<dyn CompletionProvider>, Error> {
60        match provider {
61            Providers::Ollama => {
62                let client = ollama::OllamaAI::new(
63                    model,
64                    Some(&system_message),
65                    temperature,
66                    max_tokens,
67                    tool,
68                )?;
69                Ok(Box::new(client))
70            }
71            Providers::OpenAI => {
72                let api_key = api_key.ok_or_else(|| {
73                    Error::AuthenticationError(
74                        "api_key is required for openai provider".to_string(),
75                    )
76                })?;
77
78                let client = openai::OpenAI::new(
79                    api_key,
80                    model,
81                    Some(&system_message),
82                    temperature,
83                    max_tokens,
84                    tool,
85                )?;
86                Ok(Box::new(client))
87            }
88        }
89    }
90
91    pub fn init_with_schema<J: JsonSchema, T: Tool + 'static>(
92        provider: Providers,
93        model: &str,
94        api_key: Option<&str>,
95        system_message: String,
96        temperature: Option<f64>,
97        max_tokens: Option<u64>,
98        tool: Option<ToolWrapper<T>>,
99    ) -> Result<Box<dyn CompletionProvider>, Error> {
100        match provider {
101            Providers::Ollama => {
102                let client = ollama::OllamaAI::new_with_schema::<J, T>(
103                    model,
104                    Some(&system_message),
105                    temperature,
106                    max_tokens,
107                    tool,
108                )?;
109                Ok(Box::new(client))
110            }
111            Providers::OpenAI => {
112                let api_key = api_key.ok_or_else(|| {
113                    Error::AuthenticationError(
114                        "api_key is required for openai provider".to_string(),
115                    )
116                })?;
117
118                let client = openai::OpenAI::new_with_schema::<J, T>(
119                    api_key,
120                    model,
121                    Some(&system_message),
122                    temperature,
123                    max_tokens,
124                    tool,
125                )?;
126                Ok(Box::new(client))
127            }
128        }
129    }
130}