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