Orchestration (#7999)
Signed-off-by: Douwe Osinga <douwe@squareup.com> Co-authored-by: Douwe Osinga <douwe@squareup.com> Co-authored-by: Tyler Longwell <tlongwell@squareup.com>
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
use crate::routes::errors::ErrorResponse;
|
||||
use crate::routes::reply::{get_token_state, track_tool_telemetry, MessageEvent};
|
||||
use crate::session_event_bus::RequestGuard;
|
||||
use crate::state::AppState;
|
||||
use axum::{
|
||||
extract::{DefaultBodyLimit, Path, State},
|
||||
@@ -320,7 +321,13 @@ pub async fn session_reply(
|
||||
}
|
||||
|
||||
let bus = state.get_or_create_event_bus(&session_id).await;
|
||||
let cancel_token = bus.register_request(request_id.clone()).await;
|
||||
|
||||
let cancel_token = bus
|
||||
.try_register_request(request_id.clone())
|
||||
.await
|
||||
.map_err(|_| {
|
||||
ErrorResponse::bad_request("Session already has an active request. Cancel it first.")
|
||||
})?;
|
||||
|
||||
let user_message = request.user_message;
|
||||
let override_conversation = request.override_conversation;
|
||||
@@ -332,6 +339,8 @@ pub async fn session_reply(
|
||||
let task_bus = bus.clone();
|
||||
|
||||
drop(tokio::spawn(async move {
|
||||
let mut _guard = RequestGuard::new(task_bus.clone(), task_request_id.clone());
|
||||
|
||||
let publish = |rid: Option<String>, event: MessageEvent| {
|
||||
let bus = task_bus.clone();
|
||||
async move {
|
||||
@@ -350,7 +359,6 @@ pub async fn session_reply(
|
||||
},
|
||||
)
|
||||
.await;
|
||||
task_bus.cleanup_request(&task_request_id).await;
|
||||
return;
|
||||
}
|
||||
};
|
||||
@@ -370,7 +378,6 @@ pub async fn session_reply(
|
||||
},
|
||||
)
|
||||
.await;
|
||||
task_bus.cleanup_request(&task_request_id).await;
|
||||
return;
|
||||
}
|
||||
};
|
||||
@@ -420,7 +427,6 @@ pub async fn session_reply(
|
||||
},
|
||||
)
|
||||
.await;
|
||||
task_bus.cleanup_request(&task_request_id).await;
|
||||
return;
|
||||
}
|
||||
};
|
||||
@@ -562,6 +568,7 @@ pub async fn session_reply(
|
||||
)
|
||||
.await;
|
||||
|
||||
_guard.disarm();
|
||||
task_bus.cleanup_request(&task_request_id).await;
|
||||
}));
|
||||
|
||||
|
||||
@@ -116,7 +116,7 @@ impl SessionEventBus {
|
||||
requests.keys().cloned().collect()
|
||||
}
|
||||
|
||||
/// Register a new request and return its cancellation token.
|
||||
#[cfg(test)]
|
||||
pub async fn register_request(&self, request_id: String) -> CancellationToken {
|
||||
let token = CancellationToken::new();
|
||||
let mut requests = self.active_requests.lock().await;
|
||||
@@ -124,6 +124,20 @@ impl SessionEventBus {
|
||||
token
|
||||
}
|
||||
|
||||
/// Atomically check no requests are active and register one. Returns Err if busy.
|
||||
pub async fn try_register_request(
|
||||
&self,
|
||||
request_id: String,
|
||||
) -> Result<CancellationToken, String> {
|
||||
let mut requests = self.active_requests.lock().await;
|
||||
if !requests.is_empty() {
|
||||
return Err("Session already has an active request".into());
|
||||
}
|
||||
let token = CancellationToken::new();
|
||||
requests.insert(request_id, token.clone());
|
||||
Ok(token)
|
||||
}
|
||||
|
||||
/// Cancel a specific request by request_id.
|
||||
pub async fn cancel_request(&self, request_id: &str) -> bool {
|
||||
let requests = self.active_requests.lock().await;
|
||||
@@ -156,6 +170,38 @@ impl Default for SessionEventBus {
|
||||
}
|
||||
}
|
||||
|
||||
pub struct RequestGuard {
|
||||
bus: std::sync::Arc<SessionEventBus>,
|
||||
request_id: String,
|
||||
disarmed: bool,
|
||||
}
|
||||
|
||||
impl RequestGuard {
|
||||
pub fn new(bus: std::sync::Arc<SessionEventBus>, request_id: String) -> Self {
|
||||
Self {
|
||||
bus,
|
||||
request_id,
|
||||
disarmed: false,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn disarm(&mut self) {
|
||||
self.disarmed = true;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for RequestGuard {
|
||||
fn drop(&mut self) {
|
||||
if !self.disarmed {
|
||||
let bus = self.bus.clone();
|
||||
let request_id = self.request_id.clone();
|
||||
tokio::spawn(async move {
|
||||
bus.cleanup_request(&request_id).await;
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
Reference in New Issue
Block a user