use std::{net::SocketAddr, sync::Arc};
use futures::{StreamExt as _, TryStreamExt, future};
use tarpc::{
server::{self, Channel as _, incoming::Incoming as _},
tokio_serde::formats::Json,
};
use crate::{Error, error::ApiError, providers::CompletionProvider};
#[derive(Clone)]
pub struct AgentServer {
pub socket_addr: SocketAddr,
pub providers: Arc<Box<dyn CompletionProvider>>,
}
impl AgentServer {
pub async fn run(self) -> Result<(), Error> {
let mut listener =
tarpc::serde_transport::tcp::listen(self.socket_addr, Json::default).await?;
listener.config_mut().max_frame_length(usize::MAX);
println!("Listening on: {}", listener.local_addr());
listener
.map_err(|e| eprintln!("{}", e)) .filter_map(|r| future::ready(r.ok()))
.map(server::BaseChannel::with_defaults)
.max_channels_per_key(1, |t| t.transport().peer_addr().unwrap().ip())
.map(|channel| {
channel.execute(self.clone().serve()).for_each(|f| async {
tokio::spawn(f);
})
})
.buffer_unordered(10)
.for_each(|_| async {})
.await;
Ok(())
}
}
#[tarpc::service]
trait AgentWorker {
async fn message(user_message: String) -> Result<String, ApiError>;
}
impl AgentWorker for AgentServer {
async fn message(
self,
_context: ::tarpc::context::Context,
user_message: String,
) -> Result<String, ApiError> {
println!("Message received");
self.providers
.chat(&user_message)
.await
.map_err(ApiError::from)
}
}