Skip to main content

rpc_agent/
agent.rs

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