use std::time::Duration;
use futures::stream::StreamExt;
use pyo3::prelude::*;
use tokio::sync::mpsc;
use super::{
AgentId, PyCommand,
handlers::{agent, async_ops, chat, query},
};
pub(crate) const AGY_BRIDGE_GLOBALS_MODULE: &str = "_agy_bridge_globals";
pub(crate) type RegisteredAgentPair = (std::sync::Arc<Py<PyAny>>, std::sync::Arc<Py<PyAny>>);
pub(crate) type RegistryInner = std::collections::HashMap<AgentId, RegisteredAgentPair>;
pub(crate) type AgentRegistry = std::sync::Arc<std::sync::Mutex<RegistryInner>>;
const PY_AEXIT_TIMEOUT: Duration = Duration::from_secs(3);
const AEXIT_GRACE: Duration = Duration::from_secs(2);
const LEFTOVER_AEXIT_TIMEOUT: Duration = PY_AEXIT_TIMEOUT.saturating_add(AEXIT_GRACE);
pub(in crate::runtime) type ActiveChatWriters =
std::sync::Arc<std::sync::Mutex<Vec<std::sync::Weak<crate::streaming::ChatResponseWriter>>>>;
fn lock_active_chats(
active_chats: &ActiveChatWriters,
) -> std::sync::MutexGuard<'_, Vec<std::sync::Weak<crate::streaming::ChatResponseWriter>>> {
active_chats.lock().unwrap_or_else(|e| {
tracing::warn!("Active-chat registry mutex poisoned — recovering: {e}");
e.into_inner()
})
}
pub(in crate::runtime) fn register_active_chat(
active_chats: &ActiveChatWriters,
writer: &std::sync::Arc<crate::streaming::ChatResponseWriter>,
) {
let mut guard = lock_active_chats(active_chats);
guard.retain(|weak| weak.strong_count() > 0);
guard.push(std::sync::Arc::downgrade(writer));
}
fn fail_active_chats(active_chats: &ActiveChatWriters, reason: &str) {
let writers = std::mem::take(&mut *lock_active_chats(active_chats));
for weak in writers {
let Some(writer) = weak.upgrade() else {
continue;
};
if let Err(e) = writer
.error_tx
.try_send(crate::streaming::StreamError::new(reason))
{
tracing::warn!(
error = %e,
"Could not deliver shutdown error to an in-flight chat \
(error channel full or closed)"
);
}
}
}
pub(super) fn lookup_agent_instance(
registry: &AgentRegistry,
agent_id: AgentId,
) -> Option<RegisteredAgentPair> {
let lock = registry.lock().unwrap_or_else(|e| {
tracing::warn!(
"Agent registry mutex poisoned — recovering (data is safe because entries \
are always fully formed before insertion): {e}"
);
e.into_inner()
});
lock.get(&agent_id).cloned()
}
pub(crate) async fn run_async_command_loop(
event_loop: Py<PyAny>,
mut cmd_rx: mpsc::Receiver<PyCommand>,
inter_agent_delay: Duration,
stream_limits: super::streaming::StreamLimits,
shutdown_timeout: Duration,
) -> PyResult<()> {
let registry: AgentRegistry = std::sync::Arc::new(std::sync::Mutex::new(RegistryInner::new()));
let event_loop = std::sync::Arc::new(event_loop);
let mut active_tasks =
futures::stream::FuturesUnordered::<futures::future::BoxFuture<'static, ()>>::new();
let rate_limiter = chat::ChatRateLimiter::new(inter_agent_delay);
let active_chats = ActiveChatWriters::default();
loop {
tokio::select! {
cmd_opt = cmd_rx.recv() => {
let Some(cmd) = cmd_opt else {
break;
};
tracing::debug!("Live-SDK command loop: received command");
if let DispatchResult::Shutdown = dispatch_async_command(
cmd,
®istry,
&event_loop,
&rate_limiter,
stream_limits,
&mut active_tasks,
&active_chats,
) {
break;
}
}
_ = active_tasks.next(), if !active_tasks.is_empty() => {
}
}
}
cmd_rx.close();
drain_active_tasks(&mut active_tasks, &active_chats, shutdown_timeout / 2).await;
drop(active_tasks);
cleanup_remaining_agents(®istry).await;
Ok(())
}
async fn drain_active_tasks(
active_tasks: &mut futures::stream::FuturesUnordered<futures::future::BoxFuture<'static, ()>>,
active_chats: &ActiveChatWriters,
deadline: Duration,
) {
if active_tasks.is_empty() {
return;
}
tracing::info!(
pending = active_tasks.len(),
timeout_ms = deadline.as_millis(),
"Shutdown: draining in-flight tasks before agent cleanup"
);
let drain = async {
while active_tasks.next().await.is_some() {
}
};
if tokio::time::timeout(deadline, drain).await.is_err() {
tracing::warn!(
pending = active_tasks.len(),
timeout_ms = deadline.as_millis(),
"Shutdown drain deadline expired — failing chats that are still streaming"
);
fail_active_chats(
active_chats,
"runtime shut down while the turn was still streaming",
);
}
}
async fn cleanup_remaining_agents(registry: &AgentRegistry) {
let remaining: Vec<_> = registry
.lock()
.unwrap_or_else(|e| {
tracing::warn!("Agent registry mutex poisoned during cleanup — recovering: {e}");
e.into_inner()
})
.drain()
.collect();
if !remaining.is_empty() {
tracing::info!(
count = remaining.len(),
"Cleaning up agents remaining in registry after command loop exit"
);
}
for (agent_id, (ctx_py, _instance)) in remaining {
tracing::debug!(agent_id = ?agent_id, "Calling __aexit__ on leftover agent");
if let Err(e) = tokio::time::timeout(
LEFTOVER_AEXIT_TIMEOUT,
cleanup_single_agent(agent_id, ctx_py),
)
.await
{
tracing::warn!(agent_id = ?agent_id, error = %e, "Timed out waiting for leftover agent __aexit__");
}
}
}
async fn cleanup_single_agent(agent_id: AgentId, ctx_py: std::sync::Arc<Py<PyAny>>) {
let aexit_result = Python::attach(|py| {
let ctx_bound = ctx_py.bind(py);
let none = py.None();
let coro = ctx_bound.call_method1("__aexit__", (&none, &none, &none))?;
Ok::<_, PyErr>(coro.clone().unbind())
});
match aexit_result {
Ok(aexit_coro_py) => {
let aexit_fut = Python::attach(|py| {
let coro = aexit_coro_py.into_bound(py);
pyo3_async_runtimes::tokio::into_future(coro)
});
match aexit_fut {
Ok(fut) => match fut.await {
Ok(_) => {
tracing::debug!(agent_id = ?agent_id, "Agent __aexit__ completed");
}
Err(e) => {
tracing::warn!(
agent_id = ?agent_id,
error = %e,
"Agent __aexit__ returned error during cleanup"
);
}
},
Err(e) => {
tracing::warn!(
agent_id = ?agent_id,
error = %e,
"Failed to convert __aexit__ coro to future"
);
}
}
}
Err(e) => {
tracing::warn!(
agent_id = ?agent_id,
error = %e,
"Failed to call __aexit__ during cleanup"
);
}
}
match super::bridge_state().write() {
Ok(mut map) => {
map.remove(&agent_id.0);
}
Err(e) => {
tracing::warn!(
agent_id = agent_id.0,
error = %e,
"BRIDGE_STATE RwLock poisoned during cleanup"
);
}
}
}
enum DispatchResult {
Continue,
Shutdown,
}
fn dispatch_query_command(cmd: PyCommand, registry: &AgentRegistry) -> Result<(), PyCommand> {
match cmd {
PyCommand::GetHistory { agent_id, reply } => {
query::handle_get_history(registry, agent_id, reply);
}
PyCommand::GetTurnCount { agent_id, reply } => {
query::handle_get_turn_count(registry, agent_id, reply);
}
PyCommand::GetTotalUsage { agent_id, reply } => {
query::handle_get_total_usage(registry, agent_id, reply);
}
PyCommand::GetLastTurnUsage { agent_id, reply } => {
query::handle_get_last_turn_usage(registry, agent_id, reply);
}
PyCommand::GetCompactionIndices { agent_id, reply } => {
query::handle_get_compaction_indices(registry, agent_id, reply);
}
PyCommand::GetLastResponse { agent_id, reply } => {
query::handle_get_last_response(registry, agent_id, reply);
}
PyCommand::IsIdle { agent_id, reply } => {
query::handle_is_idle(registry, agent_id, reply);
}
PyCommand::GetActiveAgentCount { reply } => {
query::handle_get_active_agent_count(registry, reply);
}
other => return Err(other),
}
Ok(())
}
fn dispatch_async_command(
cmd: PyCommand,
registry: &AgentRegistry,
event_loop: &std::sync::Arc<Py<PyAny>>,
rate_limiter: &chat::ChatRateLimiter,
stream_limits: super::streaming::StreamLimits,
active_tasks: &mut futures::stream::FuturesUnordered<futures::future::BoxFuture<'static, ()>>,
active_chats: &ActiveChatWriters,
) -> DispatchResult {
let cmd = match dispatch_query_command(cmd, registry) {
Ok(()) => return DispatchResult::Continue,
Err(cmd) => cmd,
};
let cmd = match dispatch_lifecycle_command(cmd, registry, event_loop, active_tasks) {
Ok(()) => return DispatchResult::Continue,
Err(cmd) => cmd,
};
let cmd = match cmd {
PyCommand::Chat {
agent_id,
prompt,
reply,
} => {
if let Some(task) = chat::dispatch_chat_command(
registry,
agent_id,
prompt,
reply,
rate_limiter.clone(),
stream_limits,
std::sync::Arc::clone(active_chats),
) {
active_tasks.push(task);
}
return DispatchResult::Continue;
}
other => other,
};
let cmd = match dispatch_agent_operation(cmd, registry, active_tasks) {
Ok(()) => return DispatchResult::Continue,
Err(cmd) => cmd,
};
match cmd {
PyCommand::Shutdown => {
tracing::info!("Shutdown command received, exiting async command loop");
DispatchResult::Shutdown
}
unhandled => {
tracing::error!(
"Unhandled PyCommand variant reached the final dispatch phase — dropping it. \
This is a bug: add the variant to one of the dispatch phases. The caller \
will observe Error::ChannelClosed."
);
drop(unhandled);
DispatchResult::Continue
}
}
}
fn spawn_agent_task(
active_tasks: &mut futures::stream::FuturesUnordered<futures::future::BoxFuture<'static, ()>>,
fut: impl std::future::Future<Output = ()> + Send + 'static,
) {
active_tasks.push(Box::pin(fut));
}
fn dispatch_lifecycle_command(
cmd: PyCommand,
registry: &AgentRegistry,
event_loop: &std::sync::Arc<Py<PyAny>>,
active_tasks: &mut futures::stream::FuturesUnordered<futures::future::BoxFuture<'static, ()>>,
) -> Result<(), PyCommand> {
match cmd {
PyCommand::CreateAgent {
agent_id,
config_json,
reply,
} => {
let registry = registry.clone();
let event_loop = std::sync::Arc::clone(event_loop);
spawn_agent_task(active_tasks, async move {
agent::handle_create_agent(registry, event_loop, agent_id, config_json, reply)
.await;
});
}
PyCommand::ShutdownAgent { agent_id, reply } => {
let registry = registry.clone();
spawn_agent_task(active_tasks, async move {
agent::handle_shutdown_agent(registry, agent_id, reply).await;
});
}
other => return Err(other),
}
Ok(())
}
fn dispatch_agent_operation(
cmd: PyCommand,
registry: &AgentRegistry,
active_tasks: &mut futures::stream::FuturesUnordered<futures::future::BoxFuture<'static, ()>>,
) -> Result<(), PyCommand> {
match cmd {
PyCommand::Cancel { agent_id, reply } => {
let registry = registry.clone();
spawn_agent_task(active_tasks, async move {
async_ops::handle_cancel(registry, agent_id, reply).await;
});
}
PyCommand::WaitForIdle { agent_id, reply } => {
let registry = registry.clone();
spawn_agent_task(active_tasks, async move {
async_ops::handle_wait_for_idle(registry, agent_id, reply).await;
});
}
PyCommand::ClearHistory { agent_id, reply } => {
async_ops::handle_clear_history(registry, agent_id, reply);
}
PyCommand::Send {
agent_id,
prompt,
reply,
} => {
let registry = registry.clone();
spawn_agent_task(active_tasks, async move {
async_ops::handle_send(registry, agent_id, prompt, reply).await;
});
}
PyCommand::SignalIdle { agent_id, reply } => {
let registry = registry.clone();
spawn_agent_task(active_tasks, async move {
async_ops::handle_signal_idle(registry, agent_id, reply).await;
});
}
PyCommand::WaitForWakeup {
agent_id,
timeout_secs,
reply,
} => {
let registry = registry.clone();
spawn_agent_task(active_tasks, async move {
async_ops::handle_wait_for_wakeup(registry, agent_id, timeout_secs, reply).await;
});
}
PyCommand::Delete { agent_id, reply } => {
let registry = registry.clone();
spawn_agent_task(active_tasks, async move {
async_ops::handle_delete(registry, agent_id, reply).await;
});
}
PyCommand::Disconnect { agent_id, reply } => {
let registry = registry.clone();
spawn_agent_task(active_tasks, async move {
async_ops::handle_disconnect(registry, agent_id, reply).await;
});
}
other => return Err(other),
}
Ok(())
}