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        println!("Listening on: {}", listener.local_addr());
32
33        listener
34            .map_err(|e| eprintln!("{}", e)) // TODO: Improve error handling.
35            .filter_map(|r| future::ready(r.ok()))
36            .map(server::BaseChannel::with_defaults)
37            // Limit channels to 1 per IP.
38            .max_channels_per_key(1, |t| {
39                t.transport()
40                    .peer_addr()
41                    .map(|addr| addr.ip())
42                    .unwrap_or(IpAddr::V4(Ipv4Addr::UNSPECIFIED))
43            })
44            .map(|channel| {
45                channel.execute(self.clone().serve()).for_each(|f| async {
46                    tokio::spawn(f);
47                })
48            })
49            // Max 10 channels.
50            .buffer_unordered(10)
51            .for_each(|_| async {})
52            .await;
53
54        Ok(())
55    }
56
57    #[cfg(test)]
58    pub(crate) fn new(
59        socket_addr: SocketAddr,
60        providers: Arc<Box<dyn CompletionProvider>>,
61    ) -> Self {
62        Self {
63            socket_addr,
64            providers,
65        }
66    }
67}
68
69#[tarpc::service]
70pub(crate) trait AgentWorker {
71    async fn message(user_message: Message) -> Result<String, ApiError>;
72}
73
74impl AgentWorker for AgentServer {
75    /// Handles a user message by passing it to the completion provider and returning the response.
76    async fn message(
77        self,
78        _context: ::tarpc::context::Context,
79        user_message: Message,
80    ) -> Result<String, ApiError> {
81        println!("Message received");
82        self.providers
83            .chat(&user_message.to_string())
84            .await
85            .map_err(ApiError::from)
86    }
87}