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