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 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)) .filter_map(|r| future::ready(r.ok()))
29 .map(server::BaseChannel::with_defaults)
30 .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 .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 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}