diff --git a/signal-gateway-assistant/src/lib.rs b/signal-gateway-assistant/src/lib.rs index 17d0a92..7a0c683 100644 --- a/signal-gateway-assistant/src/lib.rs +++ b/signal-gateway-assistant/src/lib.rs @@ -1,7 +1,9 @@ //! Assistant API used by signal-gateway. //! -//! This crate provides abstract types for LLM assistant interactions, -//! making it easy to swap in different LLM implementations. +//! Implement the `Assistant` trait to swap in your own LLM, context management, etc. +//! +//! The assistant trait is unopinionated and you should be able to use something like +//! `rsllm` or `rig` with relative ease if you want to. mod assistant; mod chat_message; diff --git a/signal-gateway/src/assistant/worker.rs b/signal-gateway/src/assistant/worker.rs index 90ee6df..2276328 100644 --- a/signal-gateway/src/assistant/worker.rs +++ b/signal-gateway/src/assistant/worker.rs @@ -27,7 +27,6 @@ pub struct AssistantWorker { assistant: Box, input_rx: mpsc::Receiver, stop_rx: mpsc::Receiver<()>, - cancel_token: CancellationToken, } impl AssistantWorker { @@ -41,7 +40,6 @@ impl AssistantWorker { assistant, input_rx, stop_rx, - cancel_token: CancellationToken::new(), } } @@ -65,10 +63,22 @@ impl AssistantWorker { async fn handle_input(&mut self, input: Input) { match input { Input::Prompt(msg, sender) => { - // Reset the cancel token for each new request - self.cancel_token = CancellationToken::new(); + // Make a cancel token for each new request + let cancel_token = CancellationToken::new(); - let result = self.assistant.prompt(msg, self.cancel_token.clone()).await; + let mut assistant_fut = self.assistant.prompt(msg, cancel_token.clone()); + + // If the assistant finishes normally, return its result. + // If we get a stop request, cancel the token, then wait for assistant to finish. + let result = tokio::select! { + result = &mut assistant_fut => { + result + }, + _ = self.stop_rx.recv() => { + cancel_token.cancel(); + assistant_fut.await + } + }; let response = match result { Ok(Some(resp)) => Ok(assistant_response_to_admin(resp)), @@ -95,9 +105,6 @@ impl AssistantWorker { } async fn handle_stop(&mut self) { - // Cancel any in-progress request - self.cancel_token.cancel(); - // Drain remaining stop signals while self.stop_rx.try_recv().is_ok() {}