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 #[cfg(feature = "tracing")]
32 tracing::info!("Listening on: {}", listener.local_addr());
33
34 #[cfg(not(feature = "tracing"))]
35 println!("Listening on: {}", listener.local_addr());
36
37 listener
38 .map_err(|e| eprintln!("{}", e)) .filter_map(|r| future::ready(r.ok()))
40 .map(server::BaseChannel::with_defaults)
41 .max_channels_per_key(1, |t| {
43 t.transport()
44 .peer_addr()
45 .map(|addr| addr.ip())
46 .unwrap_or(IpAddr::V4(Ipv4Addr::UNSPECIFIED))
47 })
48 .map(|channel| {
49 channel.execute(self.clone().serve()).for_each(|f| async {
50 tokio::spawn(f);
51 })
52 })
53 .buffer_unordered(10)
55 .for_each(|_| async {})
56 .await;
57
58 Ok(())
59 }
60
61 #[cfg(test)]
62 pub(crate) fn new(
63 socket_addr: SocketAddr,
64 providers: Arc<Box<dyn CompletionProvider>>,
65 ) -> Self {
66 Self {
67 socket_addr,
68 providers,
69 }
70 }
71}
72
73#[tarpc::service]
74pub(crate) trait AgentWorker {
75 async fn message(user_message: Message) -> Result<String, ApiError>;
76}
77
78impl AgentWorker for AgentServer {
79 #[cfg_attr(
81 feature = "tracing",
82 tracing::instrument(name = "agent.message", skip(self, _context, user_message))
83 )]
84 async fn message(
85 self,
86 _context: ::tarpc::context::Context,
87 user_message: Message,
88 ) -> Result<String, ApiError> {
89 let prompt: String = user_message.try_into()?;
90 self.providers.chat(&prompt).await.map_err(ApiError::from)
91 }
92}