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:
Douwe Osinga
2026-03-19 16:30:24 -04:00
committed by GitHub
parent 53236ff327
commit 3f8eec3237
8 changed files with 794 additions and 11 deletions
@@ -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;
}));
+47 -1
View File
@@ -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::*;