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}
15
16pub enum Providers {
18 Ollama,
20 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}