Skip to main content

rpc_agent/
agent.rs

1use std::{
2    net::{IpAddr, Ipv4Addr, SocketAddr},
3    sync::Arc,
4};
5
6use futures::{StreamExt as _, TryStreamExt, future};
7use tarpc::{
8    server::{self, Channel as _, incoming::Incoming as _},
9    tokio_serde::formats::Json,
10};
11
12use crate::{
13    error::{ApiError, Error},
14    message::Message,
15    providers::CompletionProvider,
16};
17
18#[derive(Clone)]
19pub struct AgentServer {
20    pub(crate) socket_addr: SocketAddr,
21    pub(crate) providers: Arc<Box<dyn CompletionProvider>>,
22}
23
24impl AgentServer {
25    /// Runs the agent server.
26    pub async fn run(self) -> Result<(), Error> {
27        let mut listener =
28            tarpc::serde_transport::tcp::listen(self.socket_addr, Json::default).await?;
29        listener.config_mut().max_frame_length(usize::MAX);
30
31        #[cfg(feature = "tracing")]
32        tracing::info!("Listening on: {}", listener.local_addr());
33
34        #[cfg(not(feature = "tracing"))]
35        println!("Listening on: {}", listener.local_addr());
36
37        listener
38            .map_err(|e| eprintln!("{}", e)) // TODO: Improve error handling.
39            .filter_map(|r| future::ready(r.ok()))
40            .map(server::BaseChannel::with_defaults)
41            // Limit channels to 1 per IP.
42            .max_channels_per_key(1, |t| {
43                t.transport()
44                    .peer_addr()
45                    .map(|addr| addr.ip())
46                    .unwrap_or(IpAddr::V4(Ipv4Addr::UNSPECIFIED))
47            })
48            .map(|channel| {
49                channel.execute(self.clone().serve()).for_each(|f| async {
50                    tokio::spawn(f);
51                })
52            })
53            // Max 10 channels.
54            .buffer_unordered(10)
55            .for_each(|_| async {})
56            .await;
57
58        Ok(())
59    }
60
61    #[cfg(test)]
62    pub(crate) fn new(
63        socket_addr: SocketAddr,
64        providers: Arc<Box<dyn CompletionProvider>>,
65    ) -> Self {
66        Self {
67            socket_addr,
68            providers,
69        }
70    }
71}
72
73#[tarpc::service]
74pub(crate) trait AgentWorker {
75    async fn message(user_message: Message) -> Result<String, ApiError>;
76}
77
78impl AgentWorker for AgentServer {
79    /// Handles a user message by passing it to the completion provider and returning the response.
80    #[cfg_attr(
81        feature = "tracing",
82        tracing::instrument(name = "agent.message", skip(self, _context, user_message))
83    )]
84    async fn message(
85        self,
86        _context: ::tarpc::context::Context,
87        user_message: Message,
88    ) -> Result<String, ApiError> {
89        let prompt: String = user_message.try_into()?;
90        self.providers.chat(&prompt).await.map_err(ApiError::from)
91    }
92}