rpc_agent/providers/
mod.rs1use 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
22pub 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}