Skip to main content

rpc_agent/
builder.rs

1use std::{net::SocketAddr, sync::Arc};
2
3use rig::tool::Tool;
4use schemars::JsonSchema;
5
6use crate::{
7    AgentServer, Providers,
8    error::Error,
9    tools::{NoTool, ToolWrapper},
10};
11
12/// Builder for creating an [`AgentServer`].
13pub struct AgentServerBuilder<'a> {
14    port: u16,
15    provider: Providers,
16    system_message: &'a str,
17    model: &'a str,
18    api_key: Option<&'a str>,
19    temperature: Option<f64>,
20    max_tokens: Option<u64>,
21}
22
23impl<'a> AgentServerBuilder<'a> {
24    /// Creates a new [`AgentServerBuilder`] with the given port, provider, system message, and model.
25    pub fn new(port: u16, provider: Providers, system_message: &'a str, model: &'a str) -> Self {
26        Self {
27            port,
28            provider,
29            system_message,
30            model,
31            api_key: None,
32            temperature: None,
33            max_tokens: None,
34        }
35    }
36
37    /// Sets the API key for the provider.
38    #[inline]
39    pub fn api_key(mut self, api_key: &'a str) -> Self {
40        self.api_key = Some(api_key);
41        self
42    }
43
44    /// Sets the temperature for the provider.
45    #[inline]
46    pub fn temperature(mut self, temperature: f64) -> Self {
47        self.temperature = Some(temperature);
48        self
49    }
50
51    /// Sets the maximum number of tokens for the provider.
52    #[inline]
53    pub fn max_tokens(mut self, max_tokens: u64) -> Self {
54        self.max_tokens = Some(max_tokens);
55        self
56    }
57
58    /// Builds the [`AgentServer`] with the given configuration.
59    pub fn build(self) -> Result<AgentServer, Error> {
60        let providers = Providers::init::<NoTool>(
61            self.provider,
62            self.model,
63            self.api_key,
64            self.system_message.to_string(),
65            self.temperature,
66            self.max_tokens,
67            None,
68        )?;
69
70        Ok(AgentServer {
71            socket_addr: SocketAddr::from(([0, 0, 0, 0], self.port)),
72            providers: Arc::new(providers),
73        })
74    }
75
76    /// Builds the [`AgentServer`] with the given configuration and schema.
77    pub fn build_with_schema<J: JsonSchema>(self) -> Result<AgentServer, Error> {
78        let providers = Providers::init_with_schema::<J, NoTool>(
79            self.provider,
80            self.model,
81            self.api_key,
82            self.system_message.to_string(),
83            self.temperature,
84            self.max_tokens,
85            None,
86        )?;
87
88        Ok(AgentServer {
89            socket_addr: SocketAddr::from(([0, 0, 0, 0], self.port)),
90            providers: Arc::new(providers),
91        })
92    }
93
94    /// Builds the [`AgentServer`] with the given configuration and tool.
95    pub fn build_with_tool<T: Tool + 'static>(
96        self,
97        tool: ToolWrapper<T>,
98    ) -> Result<AgentServer, Error> {
99        let providers = Providers::init::<T>(
100            self.provider,
101            self.model,
102            self.api_key,
103            self.system_message.to_string(),
104            self.temperature,
105            self.max_tokens,
106            Some(tool),
107        )?;
108
109        Ok(AgentServer {
110            socket_addr: SocketAddr::from(([0, 0, 0, 0], self.port)),
111            providers: Arc::new(providers),
112        })
113    }
114}