use std::collections::{HashMap, HashSet, VecDeque};
use std::fmt::Write as _;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex as StdMutex, Weak};
use std::time::Duration;
use serde_json::{Map, Value, json};
use sha2::{Digest, Sha256};
use tokio::sync::{Mutex, Notify, broadcast, watch};
use tokio::task::JoinHandle;
use tokio::time::{Instant, timeout_at};
use crate::domain::chat::{
ChatContent, ChatDeltaKind, ChatEvent, ChatExtensionKeyedRecord, ChatExtensionState,
ChatExtensionUiRecord, ChatMessage, ChatModel, ChatRole, ChatSessionStarted, ChatSessionState,
ChatSnapshot, ChatStatus, ChatStopReason, ChatToolExecution, ChatToolStatus,
ChatTranscriptPage, MAX_TRANSCRIPT_DELIVERY_BYTES, PromptAccepted, SequencedChatEvent,
TranscriptDirection,
};
use crate::domain::errors::{AgentError, AgentResult, ErrorCode};
use crate::domain::pi_rpc::{
ExtensionUiResponsePayload, PiCommandInfo, PiMessagesData, PiRpcCommand, PiRpcEvent,
PiRpcResponse, PiSessionState, PiSessionStats, ThinkingLevel,
};
use crate::domain::protocol::AgentMessage;
use crate::infrastructure::pi_history::{
HistoryPageBoundary, ParsedHistoryPage, PiHistoryStore, ResolvedPiSession,
};
use crate::infrastructure::pi_rpc_process::{
PiProcessOptions, PiRpcEventReceiver, PiRpcProcess, PiSessionMode,
};
use crate::operational::connection_routes::AuthenticatedRequestOwner;
const EVENT_CHANNEL_CAPACITY: usize = 64;
const MAX_FANOUT_EVENT_BYTES: usize = 256 * 1024;
const MAX_GATED_PUBLICATION_BYTES: usize = 16 * 1024 * 1024;
const MAX_SNAPSHOT_MESSAGE_ATTEMPTS: usize = 3;
const MAX_STABLE_MESSAGE_IDS: usize = 4096;
const MAX_PROMPT_CANCELLATION_INTENTS: usize = 64;
const MAX_ACTIVE_SESSION_IDS: usize = 1024;
const MAX_PENDING_EXTENSION_REQUESTS: usize = 32;
const MAX_EXTENSION_NOTIFICATIONS: usize = 32;
const MAX_EXTENSION_STATUSES: usize = 32;
const MAX_EXTENSION_WIDGETS: usize = 32;
const MAX_STATS_REQUEST_TIMEOUT: Duration = Duration::from_secs(1);
const MIN_SYNCHRONOUS_EXTENSION_PROMPT_TIMEOUT: Duration = Duration::from_secs(60);
const MAX_TRANSCRIPT_CURSORS_PER_SESSION: usize = 32;
const MAX_RECENT_LIVE_MESSAGES: usize = 512;
const MAX_PERSISTED_ALIASES: usize = 1024;
pub(crate) struct CreatePiChat {
pub(crate) session_id: String,
pub(crate) cwd: PathBuf,
pub(crate) provider: Option<String>,
pub(crate) model_id: Option<String>,
pub(crate) thinking_level: Option<ThinkingLevel>,
pub(crate) request_id: String,
pub(crate) prompt: String,
}
pub(crate) struct ResumePiChat {
pub(crate) session_id: String,
pub(crate) cwd: PathBuf,
pub(crate) session_path: PathBuf,
pub(crate) history_store: Option<PiHistoryStore>,
pub(crate) history_session: Option<ResolvedPiSession>,
}
pub(crate) struct ExtensionResponse {
pub(crate) id: String,
pub(crate) response: ExtensionUiResponsePayload,
}
pub(crate) struct PiRpcManager {
inner: Arc<ManagerInner>,
}
pub(crate) struct ManagerInner {
command: Vec<String>,
request_timeout: Duration,
registry: StdMutex<SessionRegistry>,
events: broadcast::Sender<SequencedChatEvent>,
next_internal_id: AtomicU64,
shutting_down: AtomicBool,
#[cfg(test)]
pump_delay_ms: AtomicU64,
#[cfg(test)]
snapshot_build_delay_ms: AtomicU64,
#[cfg(test)]
normalize_pause_enabled: AtomicBool,
#[cfg(test)]
normalize_pause_claimed: AtomicBool,
#[cfg(test)]
normalize_paused: AtomicBool,
#[cfg(test)]
normalize_paused_notify: Notify,
#[cfg(test)]
normalize_resume_notify: Notify,
#[cfg(test)]
activation_delay_ms: AtomicU64,
#[cfg(test)]
activation_ready: AtomicBool,
#[cfg(test)]
activation_ready_notify: Notify,
#[cfg(test)]
stale_cancellation_expiry_pause: TestPause,
#[cfg(test)]
prompt_pre_dispatch_pause: TestPause,
}
#[cfg(test)]
struct TestPause {
enabled: AtomicBool,
paused: AtomicBool,
paused_notify: Notify,
resume_notify: Notify,
}
#[cfg(test)]
impl TestPause {
fn new() -> Self {
Self {
enabled: AtomicBool::new(false),
paused: AtomicBool::new(false),
paused_notify: Notify::new(),
resume_notify: Notify::new(),
}
}
fn set_enabled(&self, enabled: bool) {
self.paused.store(false, Ordering::Release);
self.enabled.store(enabled, Ordering::Release);
if !enabled {
self.resume_notify.notify_waiters();
}
}
async fn wait_if_enabled(&self) {
if !self.enabled.load(Ordering::Acquire) {
return;
}
self.paused.store(true, Ordering::Release);
self.paused_notify.notify_waiters();
loop {
let notified = self.resume_notify.notified();
if !self.enabled.load(Ordering::Acquire) {
return;
}
notified.await;
}
}
async fn wait_until_paused(&self) {
loop {
let notified = self.paused_notify.notified();
if self.paused.load(Ordering::Acquire) {
return;
}
notified.await;
}
}
}
struct SessionRegistry {
sessions: HashMap<String, Arc<ManagedSession>>,
next_generation: u64,
}
struct ManagedSession {
id: String,
generation: u64,
lifecycle: StdMutex<ManagedLifecycle>,
cancel_tx: watch::Sender<bool>,
active: AtomicBool,
history_path: StdMutex<Option<PathBuf>>,
}
struct ManagedLifecycle {
session: Option<Arc<PiChatSession>>,
pump: Option<JoinHandle<()>>,
closed: bool,
}
struct InitializationGuard {
inner: Weak<ManagerInner>,
entry: Arc<ManagedSession>,
armed: bool,
}
struct ActivePromptGuard {
inner: Weak<ManagerInner>,
entry: Arc<ManagedSession>,
armed: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PromptKind {
Unknown,
AgentTurn,
SynchronousExtension,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PromptPhase {
Preflight { dispatched: bool },
Running,
Canceling,
}
#[derive(Debug, Clone)]
struct ActivePromptLifecycle {
request_id: String,
user_message_id: String,
kind: PromptKind,
phase: PromptPhase,
canceled_before_start: bool,
abort_sent: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PromptDispatch {
SendToPi,
CanceledBeforeStart,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PromptCancellation {
Wait,
SendAbortToPi,
}
struct PiChatSession {
id: String,
workspace: PathBuf,
process: Arc<PiRpcProcess>,
cache: Mutex<SessionCache>,
closing: AtomicBool,
processed_event_sequence: AtomicU64,
processed_event_notify: Notify,
prompt_lifecycle_notify: Notify,
publication_gate: Arc<Mutex<()>>,
publication_state: StdMutex<PublicationState>,
history: Option<PersistedTranscript>,
owner: Mutex<OwnerLease>,
}
pub(crate) struct OwnedWorkspace {
session_id: String,
workspace: PathBuf,
entry: Arc<ManagedSession>,
session: Arc<PiChatSession>,
}
impl OwnedWorkspace {
pub(crate) fn path(&self) -> &Path {
&self.workspace
}
}
struct OwnerLease {
owner: Option<AuthenticatedRequestOwner>,
epoch: u64,
}
#[derive(Clone)]
struct PersistedTranscript {
store: PiHistoryStore,
session: ResolvedPiSession,
}
#[derive(Clone)]
struct TranscriptCursorRecord {
session_id: String,
owner: AuthenticatedRequestOwner,
generation: String,
boundary: HistoryPageBoundary,
}
#[derive(Default)]
struct TranscriptCursorRegistry {
records: HashMap<String, TranscriptCursorRecord>,
insertion_order: VecDeque<String>,
}
enum PublicationState {
Gated {
events: VecDeque<SequencedChatEvent>,
bytes: usize,
overflowed: bool,
},
Open,
Discarded,
}
struct EventPumpContext {
manager: Weak<ManagerInner>,
generation: u64,
session: Arc<PiChatSession>,
sender: broadcast::Sender<SequencedChatEvent>,
request_timeout: Duration,
failure_sequence: u64,
pump_delay: Duration,
}
struct SessionCache {
state: ChatSessionState,
prompt_pending: bool,
active_prompt: Option<ActivePromptLifecycle>,
cancellation_intents: HashSet<String>,
cancellation_intent_order: VecDeque<String>,
canceled_message_ids: VecDeque<String>,
next_message_id: u64,
active_message_id: Option<String>,
last_assistant_message: Option<ChatMessage>,
pi_message_ids: HashMap<String, String>,
snapshot_message_ids: HashMap<String, String>,
tools: HashMap<String, ChatToolExecution>,
terminal_error_emitted: bool,
last_event_sequence: u64,
last_observed_event_sequence: u64,
live_messages: HashMap<String, ChatMessage>,
recent_live_messages: VecDeque<ChatMessage>,
persisted_aliases: HashMap<String, String>,
persisted_alias_order: VecDeque<String>,
transcript_generation: String,
transcript_cursors: TranscriptCursorRegistry,
pending_user_messages: VecDeque<(String, String)>,
terminal_status: Option<ChatStatus>,
pending_extension_requests: VecDeque<PendingExtensionRequest>,
extension_notifications: VecDeque<ChatExtensionUiRecord>,
extension_statuses: VecDeque<ChatExtensionKeyedRecord>,
extension_widgets: VecDeque<ChatExtensionKeyedRecord>,
extension_editor_text: Option<String>,
extension_title: Option<String>,
}
struct PendingExtensionRequest {
record: ChatExtensionUiRecord,
deadline: Option<Instant>,
}
impl std::ops::Deref for PiRpcManager {
type Target = ManagerInner;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
impl PiRpcManager {
pub(crate) fn new(command: Vec<String>, request_timeout: Duration) -> Self {
let (events, _) = broadcast::channel(EVENT_CHANNEL_CAPACITY);
Self {
inner: Arc::new(ManagerInner {
command,
request_timeout,
registry: StdMutex::new(SessionRegistry {
sessions: HashMap::new(),
next_generation: 1,
}),
events,
next_internal_id: AtomicU64::new(1),
shutting_down: AtomicBool::new(false),
#[cfg(test)]
pump_delay_ms: AtomicU64::new(0),
#[cfg(test)]
snapshot_build_delay_ms: AtomicU64::new(0),
#[cfg(test)]
normalize_pause_enabled: AtomicBool::new(false),
#[cfg(test)]
normalize_pause_claimed: AtomicBool::new(false),
#[cfg(test)]
normalize_paused: AtomicBool::new(false),
#[cfg(test)]
normalize_paused_notify: Notify::new(),
#[cfg(test)]
normalize_resume_notify: Notify::new(),
#[cfg(test)]
activation_delay_ms: AtomicU64::new(0),
#[cfg(test)]
activation_ready: AtomicBool::new(false),
#[cfg(test)]
activation_ready_notify: Notify::new(),
#[cfg(test)]
stale_cancellation_expiry_pause: TestPause::new(),
#[cfg(test)]
prompt_pre_dispatch_pause: TestPause::new(),
}),
}
}
pub(crate) fn subscribe(&self) -> broadcast::Receiver<SequencedChatEvent> {
self.events.subscribe()
}
#[cfg(test)]
pub(crate) fn is_history_path_active(&self, path: &Path) -> bool {
self.active_session_id_for_history_path(path).is_some()
}
pub(crate) fn active_session_id_for_history_path(&self, path: &Path) -> Option<String> {
let candidate = path.canonicalize().unwrap_or_else(|_| path.to_path_buf());
lock_std(&self.registry)
.sessions
.values()
.find_map(|entry| {
entry
.history_path()
.and_then(|active| (active == candidate).then(|| entry.id.clone()))
})
}
pub(crate) fn active_session_ids(&self) -> Vec<String> {
let mut ids = lock_std(&self.registry)
.sessions
.values()
.filter(|entry| entry.active.load(Ordering::Acquire))
.map(|entry| entry.id.clone())
.take(MAX_ACTIVE_SESSION_IDS)
.collect::<Vec<_>>();
ids.sort();
ids
}
pub(crate) async fn available_models(&self) -> AgentResult<Option<Vec<Value>>> {
let selected = {
let registry = lock_std(&self.registry);
registry
.sessions
.values()
.filter(|entry| entry.active.load(Ordering::Acquire))
.filter_map(|entry| {
entry
.session()
.map(|session| (entry.id.clone(), session.process.clone()))
})
.min_by(|left, right| left.0.cmp(&right.0))
};
let Some((session_id, process)) = selected else {
return Ok(None);
};
let response = process
.request(
PiRpcCommand::GetAvailableModels {
id: self.internal_id(&session_id, "available-models"),
},
self.request_timeout,
)
.await?;
match response {
PiRpcResponse::GetAvailableModels { data, .. } => Ok(Some(data.models)),
response => Err(unexpected_response(response, "get_available_models")),
}
}
pub(crate) async fn get_commands(
&self,
workspace: &Path,
) -> AgentResult<Option<Vec<PiCommandInfo>>> {
let workspace = workspace
.canonicalize()
.map_err(|_| AgentError::new(ErrorCode::SkillScopeDenied, "skill workspace denied"))?;
let selected = {
let registry = lock_std(&self.registry);
registry
.sessions
.values()
.filter(|entry| entry.active.load(Ordering::Acquire))
.filter_map(|entry| {
let session = entry.session()?;
let active_workspace = session
.workspace
.canonicalize()
.unwrap_or_else(|_| session.workspace.clone());
(active_workspace == workspace)
.then_some((entry.id.clone(), session.process.clone()))
})
.min_by(|left, right| left.0.cmp(&right.0))
};
let Some((session_id, process)) = selected else {
return Ok(None);
};
let response = process
.request(
PiRpcCommand::GetCommands {
id: self.internal_id(&session_id, "commands"),
},
self.request_timeout,
)
.await?;
match response {
PiRpcResponse::GetCommands { data, .. } => Ok(Some(data.commands)),
response => Err(unexpected_response(response, "get_commands")),
}
}
pub(crate) async fn workspace_for_owner(
&self,
session_id: &str,
owner: &AuthenticatedRequestOwner,
) -> AgentResult<OwnedWorkspace> {
let (entry, session) = self.active_session(session_id)?;
self.ensure_active_owned_session(session_id, &entry, &session, owner)
.await?;
Ok(OwnedWorkspace {
session_id: session_id.into(),
workspace: session.workspace.clone(),
entry,
session,
})
}
pub(crate) async fn ensure_workspace_for_owner(
&self,
workspace: &OwnedWorkspace,
owner: &AuthenticatedRequestOwner,
) -> AgentResult<()> {
self.ensure_active_owned_session(
&workspace.session_id,
&workspace.entry,
&workspace.session,
owner,
)
.await
}
pub(crate) async fn set_model_for_owner(
&self,
session_id: &str,
provider: &str,
model_id: &str,
owner: &AuthenticatedRequestOwner,
) -> AgentResult<()> {
validate_model_selection(Some(provider), Some(model_id))?;
let (entry, session) = self.active_session(session_id)?;
self.ensure_active_owned_session(session_id, &entry, &session, owner)
.await?;
session.ensure_prompt_is_idle().await?;
let liveness = owner.liveness_token();
let response = tokio::select! {
_ = liveness.cancelled() => return Err(AgentError::new(
ErrorCode::InvalidMessage,
"authenticated chat requester is no longer connected",
)),
response = session.process.request(
PiRpcCommand::SetModel {
id: self.internal_id(session_id, "set-model"),
provider: provider.to_owned(),
model_id: model_id.to_owned(),
},
self.request_timeout,
) => response,
}?;
self.ensure_active_owned_session(session_id, &entry, &session, owner)
.await?;
let model_value = expect_set_model(response)?;
session
.set_model(model_from_value(model_value, Some((provider, model_id))))
.await;
Ok(())
}
pub(crate) async fn set_thinking_level_for_owner(
&self,
session_id: &str,
level: ThinkingLevel,
owner: &AuthenticatedRequestOwner,
) -> AgentResult<()> {
let (entry, session) = self.active_session(session_id)?;
self.ensure_active_owned_session(session_id, &entry, &session, owner)
.await?;
session.ensure_prompt_is_idle().await?;
let liveness = owner.liveness_token();
let response = tokio::select! {
_ = liveness.cancelled() => return Err(AgentError::new(
ErrorCode::InvalidMessage,
"authenticated chat requester is no longer connected",
)),
response = session.process.request(
PiRpcCommand::SetThinkingLevel {
id: self.internal_id(session_id, "set-thinking"),
level,
},
self.request_timeout,
) => response,
}?;
self.ensure_active_owned_session(session_id, &entry, &session, owner)
.await?;
expect_command(response, "set_thinking_level")?;
session.set_thinking(level).await;
Ok(())
}
pub(crate) async fn commands_for_owner(
&self,
session_id: &str,
owner: &AuthenticatedRequestOwner,
) -> AgentResult<Vec<PiCommandInfo>> {
let (entry, session) = self.active_session(session_id)?;
self.ensure_active_owned_session(session_id, &entry, &session, owner)
.await?;
let liveness = owner.liveness_token();
let response = tokio::select! {
_ = liveness.cancelled() => return Err(AgentError::new(
ErrorCode::InvalidMessage,
"authenticated chat requester is no longer connected",
)),
response = session.process.request(
PiRpcCommand::GetCommands { id: self.internal_id(session_id, "commands") },
self.request_timeout,
) => response,
};
self.ensure_active_owned_session(session_id, &entry, &session, owner)
.await?;
let response = response.map_err(command_discovery_failed)?;
let commands = match response {
PiRpcResponse::GetCommands { data, .. } => Ok(data
.commands
.into_iter()
.filter(|command| {
matches!(command.source.as_str(), "extension" | "prompt" | "skill")
})
.collect()),
response => Err(command_discovery_failed(unexpected_response(
response,
"get_commands",
))),
}?;
self.ensure_active_owned_session(session_id, &entry, &session, owner)
.await?;
Ok(commands)
}
#[cfg(test)]
pub(crate) async fn create(&self, request: CreatePiChat) -> AgentResult<ChatSessionStarted> {
self.create_for_owner(request, internal_request_owner())
.await
}
pub(crate) async fn create_for_owner(
&self,
request: CreatePiChat,
owner: AuthenticatedRequestOwner,
) -> AgentResult<ChatSessionStarted> {
validate_model_selection(request.provider.as_deref(), request.model_id.as_deref())?;
owner.ensure_live()?;
let session_id = request.session_id.clone();
let entry = self.reserve(&session_id, None)?;
let mut guard = InitializationGuard::new(&self.inner, entry.clone());
if let Err(error) = owner.ensure_live() {
self.inner.remove_matching(&session_id, entry.generation);
cleanup_managed(entry).await;
guard.disarm();
return Err(error);
}
let result = self.create_reserved(request, &entry, owner).await;
match result {
Ok((started, session, publication_guard)) => {
self.wait_for_test_activation_delay().await;
if let Err(error) = session.ensure_owner_live().await {
self.inner.remove_matching(&session_id, entry.generation);
cleanup_managed(entry).await;
guard.disarm();
return Err(error);
}
if let Err(error) = self.activate_entry(&entry, || {
session.open_publications(
&self.events,
vec![
SequencedChatEvent {
event_sequence: None,
event: ChatEvent::SessionStarted {
session_id: session_id.clone(),
started: started.clone(),
},
},
SequencedChatEvent {
event_sequence: None,
event: ChatEvent::PromptAccepted {
session_id: session_id.clone(),
request_id: started.prompt.request_id.clone(),
user_message_id: started.prompt.user_message_id.clone(),
},
},
],
);
}) {
cleanup_managed(entry).await;
guard.disarm();
return Err(error);
}
drop(publication_guard);
guard.disarm();
Ok(started)
}
Err(error) => {
self.inner.remove_matching(&session_id, entry.generation);
cleanup_managed(entry).await;
guard.disarm();
Err(error)
}
}
}
async fn create_reserved(
&self,
request: CreatePiChat,
entry: &Arc<ManagedSession>,
owner: AuthenticatedRequestOwner,
) -> AgentResult<(
ChatSessionStarted,
Arc<PiChatSession>,
tokio::sync::OwnedMutexGuard<()>,
)> {
let session_id = request.session_id.clone();
let spawned = PiRpcProcess::spawn(PiProcessOptions {
command: self.command.clone(),
cwd: request.cwd.clone(),
session: PiSessionMode::New,
})
.await?;
if let Err(error) = owner.ensure_live() {
let _ = spawned.process.shutdown().await;
return Err(error);
}
let process = Arc::new(spawned.process);
let session = Arc::new(PiChatSession::starting(
session_id.clone(),
request.cwd,
process.clone(),
None,
owner.clone(),
));
let publication_guard = session.publication_gate.clone().lock_owned().await;
if let Err(error) = owner.ensure_live() {
let _ = process.shutdown().await;
return Err(error);
}
entry.attach(session.clone())?;
entry.install_pump(|| self.start_event_pump(entry, session.clone(), spawned.events))?;
entry.ensure_open()?;
let state = self
.request_state_for_owner(&session, &owner, "create-state")
.await?;
session.ensure_owner_live().await?;
entry.ensure_open()?;
entry.set_history_path(
state
.0
.session_file
.as_deref()
.map(|path| normalized_history_path(&session.workspace, path)),
);
session.update_from_pi_state(&state.0, state.1).await;
if let (Some(provider), Some(model_id)) =
(request.provider.clone(), request.model_id.clone())
{
let response = process
.request(
PiRpcCommand::SetModel {
id: self.internal_id(&session_id, "set-model"),
provider: provider.clone(),
model_id: model_id.clone(),
},
self.request_timeout,
)
.await?;
session.ensure_owner_live().await?;
entry.ensure_open()?;
let model_value = expect_set_model(response)?;
session
.set_model(model_from_value(model_value, Some((&provider, &model_id))))
.await;
}
if let Some(level) = request.thinking_level {
expect_command(
process
.request(
PiRpcCommand::SetThinkingLevel {
id: self.internal_id(&session_id, "set-thinking"),
level,
},
self.request_timeout,
)
.await?,
"set_thinking_level",
)?;
session.ensure_owner_live().await?;
entry.ensure_open()?;
session.set_thinking(level).await;
}
let prompt = self
.prepare_prompt(&session, &request.request_id, &request.prompt)
.await;
let kind = self.classify_prompt_kind(&session, &request.prompt).await;
let prompt_timeout = self.prompt_timeout_for(kind);
session.set_prompt_kind(&prompt.request_id, kind).await;
let prompt_result = process
.request_with_barrier(
PiRpcCommand::Prompt {
id: self.internal_id(&session_id, "prompt"),
message: request.prompt.clone(),
images: None,
streaming_behavior: None,
},
prompt_timeout,
)
.await
.and_then(|reply| {
expect_command(reply.response, "prompt").map(|_| reply.event_barrier)
});
let event_barrier = match prompt_result {
Ok(event_barrier) => event_barrier,
Err(error) => {
session.reject_prompt(&prompt).await;
entry.mark_closed();
return Err(error);
}
};
if kind != PromptKind::SynchronousExtension {
session.mark_prompt_running(&prompt.request_id).await;
}
session.ensure_owner_live().await?;
entry.ensure_open()?;
session.mark_initialized().await;
if kind == PromptKind::SynchronousExtension {
session
.wait_for_event_barrier(event_barrier, prompt_timeout)
.await?;
if session.complete_synchronous_extension(&prompt).await {
if session.closing.load(Ordering::Acquire) {
return Err(session_missing());
}
session.publish_events(
&self.events,
vec![SequencedChatEvent {
event_sequence: None,
event: status_event(&session_id, ChatStatus::Settled),
}],
);
self.refresh_session_stats(&session, "initial-synchronous-extension-stats")
.await;
}
}
Ok((
ChatSessionStarted {
state: session.state().await,
pid: process.pid(),
prompt,
},
session,
publication_guard,
))
}
#[cfg(test)]
pub(crate) async fn resume(&self, request: ResumePiChat) -> AgentResult<ChatSnapshot> {
self.resume_for_owner(request, internal_request_owner())
.await
}
#[cfg(test)]
pub(crate) async fn resume_for_owner(
&self,
request: ResumePiChat,
owner: AuthenticatedRequestOwner,
) -> AgentResult<ChatSnapshot> {
self.resume_for_owner_with_request_id(request, owner, "")
.await
}
pub(crate) async fn resume_for_owner_with_request_id(
&self,
request: ResumePiChat,
owner: AuthenticatedRequestOwner,
request_id: &str,
) -> AgentResult<ChatSnapshot> {
let session_id = request.session_id.clone();
owner.ensure_live()?;
if self.active_session(&session_id).is_ok() {
return self
.snapshot_for_owner_with_request_id(&session_id, &owner, request_id)
.await;
}
let entry = self.reserve(&session_id, Some(request.session_path.clone()))?;
let mut guard = InitializationGuard::new(&self.inner, entry.clone());
if let Err(error) = owner.ensure_live() {
self.inner.remove_matching(&session_id, entry.generation);
cleanup_managed(entry).await;
guard.disarm();
return Err(error);
}
let result = self
.resume_reserved(request, &entry, owner, request_id)
.await;
match result {
Ok((snapshot, session, publication_guard)) => {
self.wait_for_test_activation_delay().await;
if let Err(error) = session.ensure_owner_live().await {
self.inner.remove_matching(&session_id, entry.generation);
cleanup_managed(entry).await;
guard.disarm();
return Err(error);
}
if let Err(error) = self.activate_entry(&entry, || {
session.open_publications(&self.events, Vec::new());
}) {
cleanup_managed(entry).await;
guard.disarm();
return Err(error);
}
drop(publication_guard);
guard.disarm();
Ok(snapshot)
}
Err(error) => {
self.inner.remove_matching(&session_id, entry.generation);
cleanup_managed(entry).await;
guard.disarm();
Err(error)
}
}
}
async fn resume_reserved(
&self,
request: ResumePiChat,
entry: &Arc<ManagedSession>,
owner: AuthenticatedRequestOwner,
request_id: &str,
) -> AgentResult<(
ChatSnapshot,
Arc<PiChatSession>,
tokio::sync::OwnedMutexGuard<()>,
)> {
let session_id = request.session_id.clone();
let history = match (
request.history_store.clone(),
request.history_session.clone(),
) {
(Some(store), Some(session)) => Some(PersistedTranscript { store, session }),
(None, None) => None,
_ => {
return Err(AgentError::new(
ErrorCode::InvalidMessage,
"persisted transcript source is incomplete",
));
}
};
let spawned = PiRpcProcess::spawn(PiProcessOptions {
command: self.command.clone(),
cwd: request.cwd.clone(),
session: PiSessionMode::Resume(request.session_path),
})
.await?;
if let Err(error) = owner.ensure_live() {
let _ = spawned.process.shutdown().await;
return Err(error);
}
let process = Arc::new(spawned.process);
let session = Arc::new(PiChatSession::starting(
session_id.clone(),
request.cwd,
process.clone(),
history,
owner.clone(),
));
let publication_guard = session.publication_gate.clone().lock_owned().await;
if let Err(error) = owner.ensure_live() {
let _ = process.shutdown().await;
return Err(error);
}
entry.attach(session.clone())?;
entry.install_pump(|| self.start_event_pump(entry, session.clone(), spawned.events))?;
entry.ensure_open()?;
let state = self
.request_state_for_owner(&session, &owner, "resume-state")
.await?;
session.ensure_owner_live().await?;
entry.ensure_open()?;
session.update_from_pi_state(&state.0, state.1).await;
self.refresh_session_stats(&session, "resume-stats").await;
session.ensure_owner_live().await?;
session.mark_initialized().await;
let owner_epoch = session.owner_epoch(&owner).await?;
let snapshot = if session.history.is_some() {
self.request_history_snapshot(&session, state.1, &owner, owner_epoch, request_id)
.await?
} else {
self.request_coherent_snapshot(&session, "resume-messages")
.await?
};
session.ensure_owner_live().await?;
entry.ensure_open()?;
Ok((snapshot, session, publication_guard))
}
pub(crate) async fn prompt(
&self,
session_id: &str,
request_id: &str,
text: &str,
) -> AgentResult<PromptAccepted> {
let (entry, session) = self.active_session(session_id)?;
let publication_guard = session.publication_gate.clone().lock_owned().await;
self.ensure_active_session(session_id, &entry, &session)?;
session.ensure_prompt_is_idle().await?;
let mut prompt_guard = ActivePromptGuard::new(&self.inner, entry);
let accepted = self.prepare_prompt(&session, request_id, text).await;
let kind = self.classify_prompt_kind(&session, text).await;
let prompt_timeout = self.prompt_timeout_for(kind);
session.set_prompt_kind(&accepted.request_id, kind).await;
#[cfg(test)]
self.wait_for_test_prompt_pre_dispatch_pause().await;
if session.closing.load(Ordering::Acquire) {
session.reject_prompt(&accepted).await;
return Err(session_missing());
}
if let Err(error) = session.begin_publications() {
session.reject_prompt(&accepted).await;
prompt_guard.cleanup().await;
return Err(error);
}
if session.begin_prompt_dispatch(&accepted.request_id).await
== PromptDispatch::CanceledBeforeStart
{
if !session.complete_preflight_cancellation(&accepted).await {
prompt_guard.disarm();
return Err(session_missing());
}
let published = self.publish_batch_for_active_session(
session_id,
&session,
vec![
SequencedChatEvent {
event_sequence: None,
event: ChatEvent::PromptAccepted {
session_id: session_id.into(),
request_id: accepted.request_id.clone(),
user_message_id: accepted.user_message_id.clone(),
},
},
SequencedChatEvent {
event_sequence: None,
event: ChatEvent::PromptCanceled {
session_id: session_id.into(),
request_id: accepted.request_id.clone(),
user_message_id: accepted.user_message_id.clone(),
},
},
SequencedChatEvent {
event_sequence: None,
event: status_event(session_id, ChatStatus::Settled),
},
],
);
drop(publication_guard);
prompt_guard.disarm();
return if published {
Ok(accepted)
} else {
Err(session_missing())
};
}
let reply = session
.process
.request_with_barrier(
PiRpcCommand::Prompt {
id: self.internal_id(session_id, "prompt"),
message: text.to_owned(),
images: None,
streaming_behavior: None,
},
prompt_timeout,
)
.await;
let event_barrier = match reply {
Ok(reply) => {
let event_barrier = reply.event_barrier;
if let Err(error) = expect_command(reply.response, "prompt") {
session.reject_prompt(&accepted).await;
prompt_guard.cleanup().await;
return Err(error);
}
event_barrier
}
Err(error) => {
session.reject_prompt(&accepted).await;
prompt_guard.cleanup().await;
return Err(error);
}
};
let accepted_published = self.publish_for_active_session(
session_id,
&session,
SequencedChatEvent {
event_sequence: None,
event: ChatEvent::PromptAccepted {
session_id: session_id.into(),
request_id: accepted.request_id.clone(),
user_message_id: accepted.user_message_id.clone(),
},
},
);
if !accepted_published {
session.reject_prompt(&accepted).await;
prompt_guard.cleanup().await;
return Err(session_missing());
}
drop(publication_guard);
prompt_guard.disarm();
if kind == PromptKind::SynchronousExtension {
session
.wait_for_event_barrier(event_barrier, prompt_timeout)
.await?;
if session.complete_synchronous_extension(&accepted).await {
if session.closing.load(Ordering::Acquire) {
return Err(session_missing());
}
session.publish_events(
&self.events,
vec![SequencedChatEvent {
event_sequence: None,
event: status_event(session_id, ChatStatus::Settled),
}],
);
self.refresh_session_stats(&session, "synchronous-extension-stats")
.await;
}
}
Ok(accepted)
}
pub(crate) async fn abort(
&self,
session_id: &str,
_request_id: &str,
prompt_request_id: Option<&str>,
) -> AgentResult<()> {
let session = self.session(session_id).await?;
let Some(prompt_request_id) = prompt_request_id else {
return self.send_abort_to_pi(session_id, &session).await;
};
let deadline = Instant::now() + self.request_timeout;
let mut observed_target = false;
let mut recorded_intent = false;
loop {
if session.abort_wait_failed().await {
session.remove_cancellation_intent(prompt_request_id).await;
return Err(session_missing());
}
let notified = session.prompt_lifecycle_notify.notified();
let lifecycle_present = session.prompt_lifecycle(prompt_request_id).await;
observed_target |= lifecycle_present;
if !lifecycle_present
&& (observed_target
|| (recorded_intent
&& !session.has_cancellation_intent(prompt_request_id).await))
{
return Ok(());
}
match session.cancel_prompt_lifecycle(prompt_request_id).await {
PromptCancellation::SendAbortToPi => {
self.send_abort_to_pi(session_id, &session).await?;
}
PromptCancellation::Wait => {
recorded_intent |= !lifecycle_present;
}
}
if session.abort_wait_failed().await {
session.remove_cancellation_intent(prompt_request_id).await;
return Err(session_missing());
}
if observed_target {
notified.await;
} else if timeout_at(deadline, notified).await.is_err() {
#[cfg(test)]
self.wait_for_test_stale_cancellation_expiry_pause().await;
if session
.expire_cancellation_intent_if_lifecycle_absent(prompt_request_id)
.await
{
observed_target = true;
continue;
}
return Err(stale_prompt_cancellation());
}
}
}
async fn send_abort_to_pi(&self, session_id: &str, session: &PiChatSession) -> AgentResult<()> {
expect_command(
session
.process
.request(
PiRpcCommand::Abort {
id: self.internal_id(session_id, "abort"),
},
self.request_timeout,
)
.await?,
"abort",
)
}
#[cfg(test)]
pub(crate) async fn snapshot(&self, session_id: &str) -> AgentResult<ChatSnapshot> {
let session = self.session(session_id).await?;
let owner = session.owner().await;
self.snapshot_for_owner(session_id, &owner).await
}
#[cfg(test)]
pub(crate) async fn snapshot_for_owner(
&self,
session_id: &str,
owner: &AuthenticatedRequestOwner,
) -> AgentResult<ChatSnapshot> {
self.snapshot_for_owner_with_request_id(session_id, owner, "")
.await
}
pub(crate) async fn snapshot_for_owner_with_request_id(
&self,
session_id: &str,
owner: &AuthenticatedRequestOwner,
request_id: &str,
) -> AgentResult<ChatSnapshot> {
let session = self.session(session_id).await?;
session.ensure_owner_live_for(owner).await?;
let publication_guard = session.publication_gate.clone().lock_owned().await;
session.ensure_owner_live_for(owner).await?;
let owner_epoch = session.claim_snapshot_owner(owner).await?;
session.ensure_owner_live_for(owner).await?;
session.begin_publications()?;
let result = async {
let state = self.request_state(&session, "snapshot-state").await?;
session.ensure_owner_live_for(owner).await?;
session.update_from_pi_state(&state.0, state.1).await;
self.refresh_session_stats(&session, "snapshot-stats").await;
session.ensure_owner_live_for(owner).await?;
let snapshot = if session.history.is_some() {
self.request_history_snapshot(&session, state.1, owner, owner_epoch, request_id)
.await?
} else {
self.request_coherent_snapshot(&session, "snapshot-messages")
.await?
};
session.ensure_owner_live_for(owner).await?;
Ok(snapshot)
}
.await;
session.open_publications(&self.events, Vec::new());
drop(publication_guard);
result
}
pub(crate) async fn extension_response(
&self,
session_id: &str,
response: ExtensionResponse,
) -> AgentResult<()> {
let session = self.session(session_id).await?;
let extension_request_id = response.id.clone();
session
.process
.send(
PiRpcCommand::ExtensionUiResponse {
id: response.id,
response: response.response,
},
self.request_timeout,
)
.await?;
session
.cache
.lock()
.await
.pending_extension_requests
.retain(|pending| pending.record.id != extension_request_id);
Ok(())
}
pub(crate) async fn close(&self, session_id: &str) -> AgentResult<()> {
let entry = self.remove_entry(session_id);
let Some(entry) = entry else {
if self.shutting_down.load(Ordering::Acquire) {
return Ok(());
}
return Err(session_missing());
};
cleanup_managed(entry).await;
Ok(())
}
pub(crate) async fn expire_transcript_owner(&self, owner: &AuthenticatedRequestOwner) {
let sessions = lock_std(&self.registry)
.sessions
.values()
.filter_map(|entry| entry.session())
.collect::<Vec<_>>();
for session in sessions {
session.release_transcript_owner(owner).await;
}
}
pub(crate) async fn shutdown(&self) {
self.shutting_down.store(true, Ordering::Release);
let sessions = {
let mut registry = lock_std(&self.registry);
registry
.sessions
.drain()
.map(|(_, entry)| entry)
.collect::<Vec<_>>()
};
for entry in sessions {
cleanup_managed(entry).await;
}
}
#[cfg(test)]
pub(crate) fn set_test_pump_delay(&self, delay: Duration) {
self.pump_delay_ms.store(
u64::try_from(delay.as_millis()).unwrap_or(u64::MAX),
Ordering::Release,
);
}
#[cfg(test)]
pub(crate) fn set_test_snapshot_build_delay(&self, delay: Duration) {
self.snapshot_build_delay_ms.store(
u64::try_from(delay.as_millis()).unwrap_or(u64::MAX),
Ordering::Release,
);
}
#[cfg(test)]
async fn wait_for_test_snapshot_build_delay(&self) {
self.wait_for_test_normalize_paused_if_enabled().await;
let delay = self.snapshot_build_delay_ms.load(Ordering::Acquire);
if delay > 0 {
tokio::time::sleep(Duration::from_millis(delay)).await;
}
}
#[cfg(not(test))]
async fn wait_for_test_snapshot_build_delay(&self) {}
#[cfg(test)]
pub(crate) fn set_test_pause_after_normalize(&self, enabled: bool) {
self.normalize_paused.store(false, Ordering::Release);
self.normalize_pause_claimed.store(false, Ordering::Release);
self.normalize_pause_enabled
.store(enabled, Ordering::Release);
if !enabled {
self.normalize_resume_notify.notify_waiters();
}
}
#[cfg(test)]
pub(crate) async fn wait_for_test_normalize_paused(&self) {
loop {
let notified = self.normalize_paused_notify.notified();
if self.normalize_paused.load(Ordering::Acquire) {
return;
}
notified.await;
}
}
#[cfg(test)]
pub(crate) fn resume_test_normalize(&self) {
self.normalize_pause_enabled.store(false, Ordering::Release);
self.normalize_resume_notify.notify_waiters();
}
#[cfg(test)]
async fn wait_for_test_normalize_paused_if_enabled(&self) {
if self.normalize_pause_enabled.load(Ordering::Acquire) {
self.wait_for_test_normalize_paused().await;
}
}
#[cfg(test)]
pub(crate) fn set_test_activation_delay(&self, delay: Duration) {
self.activation_delay_ms.store(
u64::try_from(delay.as_millis()).unwrap_or(u64::MAX),
Ordering::Release,
);
}
#[cfg(test)]
pub(crate) fn set_test_pause_before_stale_cancellation_expiry(&self, enabled: bool) {
self.stale_cancellation_expiry_pause.set_enabled(enabled);
}
#[cfg(test)]
pub(crate) async fn wait_for_test_stale_cancellation_expiry_paused(&self) {
self.stale_cancellation_expiry_pause
.wait_until_paused()
.await;
}
#[cfg(test)]
pub(crate) fn resume_test_stale_cancellation_expiry(&self) {
self.stale_cancellation_expiry_pause.set_enabled(false);
}
#[cfg(test)]
async fn wait_for_test_stale_cancellation_expiry_pause(&self) {
self.stale_cancellation_expiry_pause.wait_if_enabled().await;
}
#[cfg(test)]
pub(crate) fn set_test_pause_after_prompt_prepare(&self, enabled: bool) {
self.prompt_pre_dispatch_pause.set_enabled(enabled);
}
#[cfg(test)]
pub(crate) async fn wait_for_test_prompt_prepare_paused(&self) {
self.prompt_pre_dispatch_pause.wait_until_paused().await;
}
#[cfg(test)]
pub(crate) fn resume_test_prompt_prepare(&self) {
self.prompt_pre_dispatch_pause.set_enabled(false);
}
#[cfg(test)]
async fn wait_for_test_prompt_pre_dispatch_pause(&self) {
self.prompt_pre_dispatch_pause.wait_if_enabled().await;
}
#[cfg(test)]
async fn wait_for_test_activation_delay(&self) {
self.activation_ready.store(true, Ordering::Release);
self.activation_ready_notify.notify_waiters();
let delay = self.activation_delay_ms.load(Ordering::Acquire);
if delay > 0 {
tokio::time::sleep(Duration::from_millis(delay)).await;
}
}
#[cfg(not(test))]
async fn wait_for_test_activation_delay(&self) {}
#[cfg(test)]
pub(crate) async fn wait_for_test_activation_ready(&self) {
loop {
let notified = self.activation_ready_notify.notified();
if self.activation_ready.load(Ordering::Acquire) {
return;
}
notified.await;
}
}
#[cfg(test)]
pub(crate) async fn stderr_snapshot(&self, session_id: &str) -> AgentResult<String> {
let entry = lock_std(&self.registry)
.sessions
.get(session_id)
.cloned()
.ok_or_else(session_missing)?;
let session = entry.session().ok_or_else(session_missing)?;
Ok(session.process.stderr_snapshot().await)
}
#[cfg(test)]
pub(crate) async fn stable_identity_counts(
&self,
session_id: &str,
) -> AgentResult<(usize, usize)> {
let session = self.session(session_id).await?;
let cache = session.cache.lock().await;
Ok((cache.snapshot_message_ids.len(), cache.pi_message_ids.len()))
}
#[cfg(test)]
pub(crate) async fn seed_test_stable_mapping(
&self,
session_id: &str,
message_id: &str,
) -> AgentResult<()> {
let session = self.session(session_id).await?;
let mut cache = session.cache.lock().await;
cache.snapshot_message_ids.insert(
format!("test-canceled-snapshot:{message_id}"),
message_id.into(),
);
cache
.pi_message_ids
.insert(format!("test-canceled-pi:{message_id}"), message_id.into());
Ok(())
}
#[cfg(test)]
pub(crate) async fn promote_test_message_to_authoritative(
&self,
session_id: &str,
message_id: &str,
) -> AgentResult<()> {
let session = self.session(session_id).await?;
let mut cache = session.cache.lock().await;
cache.live_messages.remove(message_id);
cache
.recent_live_messages
.retain(|message| message.id != message_id);
cache.snapshot_message_ids.insert(
format!("test-authoritative-snapshot:{message_id}"),
message_id.into(),
);
cache.pi_message_ids.insert(
format!("test-authoritative-pi:{message_id}"),
message_id.into(),
);
Ok(())
}
#[cfg(test)]
pub(crate) async fn retain_test_message_as_pending(
&self,
session_id: &str,
message_id: &str,
) -> AgentResult<()> {
let session = self.session(session_id).await?;
let mut cache = session.cache.lock().await;
cache
.pending_user_messages
.push_back((format!("test-pending:{message_id}"), message_id.into()));
Ok(())
}
#[cfg(test)]
pub(crate) async fn canceled_message_retention_state(
&self,
session_id: &str,
message_id: &str,
) -> AgentResult<(bool, bool, bool, bool)> {
let session = self.session(session_id).await?;
let cache = session.cache.lock().await;
Ok((
cache.canceled_message_ids.contains(&message_id.to_owned()),
cache.live_messages.contains_key(message_id),
cache
.recent_live_messages
.iter()
.any(|message| message.id == message_id),
cache
.snapshot_message_ids
.values()
.chain(cache.pi_message_ids.values())
.any(|mapped| mapped == message_id),
))
}
fn reserve(
&self,
session_id: &str,
history_path: Option<PathBuf>,
) -> AgentResult<Arc<ManagedSession>> {
if self.shutting_down.load(Ordering::Acquire) {
return Err(AgentError::new(
ErrorCode::BackendDisconnected,
"Pi RPC manager is shutting down",
));
}
let mut registry = lock_std(&self.registry);
if self.shutting_down.load(Ordering::Acquire) {
return Err(AgentError::new(
ErrorCode::BackendDisconnected,
"Pi RPC manager is shutting down",
));
}
if registry.sessions.contains_key(session_id) {
return Err(session_exists());
}
let generation = registry.next_generation;
registry.next_generation = registry.next_generation.saturating_add(1);
let (cancel_tx, _) = watch::channel(false);
let entry = Arc::new(ManagedSession {
id: session_id.into(),
generation,
lifecycle: StdMutex::new(ManagedLifecycle {
session: None,
pump: None,
closed: false,
}),
cancel_tx,
active: AtomicBool::new(false),
history_path: StdMutex::new(history_path),
});
registry.sessions.insert(session_id.into(), entry.clone());
Ok(entry)
}
fn active_session(
&self,
session_id: &str,
) -> AgentResult<(Arc<ManagedSession>, Arc<PiChatSession>)> {
let entry = lock_std(&self.registry)
.sessions
.get(session_id)
.cloned()
.ok_or_else(session_missing)?;
if !entry.active.load(Ordering::Acquire) {
return Err(session_missing());
}
let session = entry.session().ok_or_else(session_missing)?;
Ok((entry, session))
}
async fn session(&self, session_id: &str) -> AgentResult<Arc<PiChatSession>> {
self.active_session(session_id).map(|(_, session)| session)
}
fn internal_id(&self, session_id: &str, purpose: &str) -> String {
let sequence = self.next_internal_id.fetch_add(1, Ordering::Relaxed);
format!("__regy:{session_id}:{purpose}:{sequence}")
}
async fn request_state(
&self,
session: &PiChatSession,
purpose: &str,
) -> AgentResult<(PiSessionState, u64)> {
let reply = session
.process
.request_with_barrier(
PiRpcCommand::GetState {
id: self.internal_id(&session.id, purpose),
},
self.request_timeout,
)
.await?;
session
.wait_for_event_barrier(reply.event_barrier, self.request_timeout)
.await?;
Ok((expect_state(reply.response)?, reply.event_barrier))
}
async fn classify_prompt_kind(&self, session: &PiChatSession, text: &str) -> PromptKind {
let Some(command_name) = slash_command_token(text) else {
return PromptKind::AgentTurn;
};
let commands = session
.process
.request(
PiRpcCommand::GetCommands {
id: self.internal_id(&session.id, "prompt-commands"),
},
self.request_timeout,
)
.await
.ok()
.and_then(|response| match response {
PiRpcResponse::GetCommands { data, .. } => Some(data.commands),
_ => None,
});
if commands.is_some_and(|commands| {
commands
.into_iter()
.any(|command| command.name == command_name && command.source == "extension")
}) {
PromptKind::SynchronousExtension
} else {
PromptKind::AgentTurn
}
}
fn prompt_timeout_for(&self, kind: PromptKind) -> Duration {
match kind {
PromptKind::SynchronousExtension => self
.request_timeout
.max(MIN_SYNCHRONOUS_EXTENSION_PROMPT_TIMEOUT),
PromptKind::Unknown | PromptKind::AgentTurn => self.request_timeout,
}
}
async fn request_state_for_owner(
&self,
session: &PiChatSession,
owner: &AuthenticatedRequestOwner,
purpose: &str,
) -> AgentResult<(PiSessionState, u64)> {
owner.ensure_live()?;
let liveness = owner.liveness_token();
let reply = tokio::select! {
_ = liveness.cancelled() => {
return Err(AgentError::new(
ErrorCode::InvalidMessage,
"authenticated chat requester is no longer connected",
));
}
reply = session.process.request_with_barrier(
PiRpcCommand::GetState {
id: self.internal_id(&session.id, purpose),
},
self.request_timeout,
) => reply?,
};
tokio::select! {
_ = liveness.cancelled() => {
return Err(AgentError::new(
ErrorCode::InvalidMessage,
"authenticated chat requester is no longer connected",
));
}
result = session.wait_for_event_barrier(reply.event_barrier, self.request_timeout) => result?,
}
owner.ensure_live()?;
Ok((expect_state(reply.response)?, reply.event_barrier))
}
async fn request_messages(
&self,
session: &PiChatSession,
purpose: &str,
) -> AgentResult<(PiMessagesData, u64)> {
let reply = session
.process
.request_with_barrier(
PiRpcCommand::GetMessages {
id: self.internal_id(&session.id, purpose),
},
self.request_timeout,
)
.await?;
session
.wait_for_event_barrier(reply.event_barrier, self.request_timeout)
.await?;
Ok((expect_messages(reply.response)?, reply.event_barrier))
}
async fn request_stats(
&self,
session: &PiChatSession,
purpose: &str,
) -> AgentResult<PiSessionStats> {
let reply = session
.process
.request_with_barrier(
PiRpcCommand::GetSessionStats {
id: self.internal_id(&session.id, purpose),
},
self.request_timeout.min(MAX_STATS_REQUEST_TIMEOUT),
)
.await?;
session
.wait_for_event_barrier(
reply.event_barrier,
self.request_timeout.min(MAX_STATS_REQUEST_TIMEOUT),
)
.await?;
expect_stats(reply.response)
}
async fn refresh_session_stats(&self, session: &PiChatSession, purpose: &str) {
let Ok(stats) = self.request_stats(session, purpose).await else {
return;
};
session.cache.lock().await.state.stats = Some(Box::new(stats.clone()));
}
fn start_event_pump(
&self,
entry: &Arc<ManagedSession>,
session: Arc<PiChatSession>,
events: PiRpcEventReceiver,
) -> JoinHandle<()> {
let sender = self.events.clone();
let cancel_rx = entry.cancel_tx.subscribe();
let weak = Arc::downgrade(&self.inner);
let generation = entry.generation;
let failure_sequence = self.next_internal_id.fetch_add(1, Ordering::Relaxed);
#[cfg(test)]
let pump_delay_ms = self.pump_delay_ms.load(Ordering::Acquire);
#[cfg(not(test))]
let pump_delay_ms = 0;
tokio::spawn(run_event_pump(
EventPumpContext {
manager: weak,
generation,
session,
sender,
request_timeout: self.request_timeout,
failure_sequence,
pump_delay: Duration::from_millis(pump_delay_ms),
},
events,
cancel_rx,
))
}
async fn prepare_prompt(
&self,
session: &PiChatSession,
request_id: &str,
text: &str,
) -> PromptAccepted {
session.prepare_prompt(request_id, text).await
}
async fn request_coherent_snapshot(
&self,
session: &PiChatSession,
purpose: &str,
) -> AgentResult<ChatSnapshot> {
for attempt in 0..MAX_SNAPSHOT_MESSAGE_ATTEMPTS {
let messages = self.request_messages(session, purpose).await?;
self.wait_for_test_snapshot_build_delay().await;
if let Some(snapshot) = self
.build_snapshot(session, messages.0.messages, messages.1)
.await
{
return Ok(snapshot);
}
if attempt + 1 == MAX_SNAPSHOT_MESSAGE_ATTEMPTS {
break;
}
}
Err(AgentError::new(
ErrorCode::SnapshotBusy,
"Pi chat changed continuously while building a coherent snapshot; retry shortly",
))
}
async fn request_history_snapshot(
&self,
session: &PiChatSession,
barrier: u64,
owner: &AuthenticatedRequestOwner,
owner_epoch: u64,
request_id: &str,
) -> AgentResult<ChatSnapshot> {
let history = session.history.clone().ok_or_else(session_missing)?;
for attempt in 0..MAX_SNAPSHOT_MESSAGE_ATTEMPTS {
let store = history.store.clone();
let resolved = history.session.clone();
let page = tokio::task::spawn_blocking(move || store.newest_page(&resolved))
.await
.map_err(|_| history_task_failed())??;
session.ensure_owner_live_for(owner).await?;
if let Some(snapshot) = self
.build_history_snapshot(session, page, barrier, owner, owner_epoch, request_id)
.await?
{
session.ensure_owner_live_for(owner).await?;
return Ok(snapshot);
}
if attempt + 1 == MAX_SNAPSHOT_MESSAGE_ATTEMPTS {
break;
}
}
Err(AgentError::new(
ErrorCode::SnapshotBusy,
"Pi chat changed continuously while building a coherent transcript snapshot; retry shortly",
))
}
#[cfg(test)]
pub(crate) async fn transcript_page(
&self,
session_id: &str,
request_id: &str,
cursor: &str,
direction: TranscriptDirection,
) -> AgentResult<AgentMessage> {
let session = self.session(session_id).await?;
let owner = session.owner().await;
self.transcript_page_for_owner(session_id, request_id, cursor, direction, &owner)
.await
}
pub(crate) async fn transcript_page_for_owner(
&self,
session_id: &str,
request_id: &str,
cursor: &str,
direction: TranscriptDirection,
owner: &AuthenticatedRequestOwner,
) -> AgentResult<AgentMessage> {
let session = self.session(session_id).await?;
session.ensure_owner_live_for(owner).await?;
let owner_epoch = session.owner_epoch(owner).await?;
let record = {
let cache = session.cache.lock().await;
let Some(record) = cache.transcript_cursors.records.get(cursor) else {
return Err(invalid_transcript_cursor());
};
if record.session_id != session_id || record.owner != *owner {
return Err(invalid_transcript_cursor());
}
if record.generation != cache.transcript_generation {
return Err(stale_transcript_cursor());
}
record.clone()
};
let history = session
.history
.clone()
.ok_or_else(invalid_transcript_cursor)?;
let store = history.store.clone();
let resolved = history.session.clone();
let boundary = record.boundary.clone();
let page =
tokio::task::spawn_blocking(move || store.page_from(&resolved, &boundary, direction))
.await
.map_err(|_| history_task_failed())?
.map_err(map_transcript_page_error)?;
session.ensure_owner_live_for(owner).await?;
let event_cursor = session.processed_event_sequence.load(Ordering::Acquire);
let transcript = {
let lease = session.owner.lock().await;
if !owner.is_live() {
drop(lease);
session.ensure_owner_live_for(owner).await?;
return Err(invalid_transcript_cursor());
}
if lease.owner.as_ref() != Some(owner) || lease.epoch != owner_epoch {
return Err(invalid_transcript_cursor());
}
let mut cache = session.cache.lock().await;
if cache.transcript_generation != record.generation {
return Err(stale_transcript_cursor());
}
transcript_page_from_history(&mut cache, session_id, owner, page, event_cursor)
};
finalize_transcript_page(session_id.to_owned(), request_id.to_owned(), transcript)
}
pub(crate) async fn finalize_snapshot_response(
&self,
session_id: &str,
request_id: &str,
snapshot: ChatSnapshot,
) -> AgentResult<AgentMessage> {
let _ = self.session(session_id).await?;
finalize_chat_snapshot(session_id.to_owned(), request_id.to_owned(), snapshot)
}
async fn build_snapshot(
&self,
session: &PiChatSession,
messages: Vec<Value>,
barrier: u64,
) -> Option<ChatSnapshot> {
let mut cache = session.cache.lock().await;
if cache.last_observed_event_sequence > barrier {
return None;
}
let snapshot_time = Instant::now();
cache.pending_extension_requests.retain(|pending| {
pending
.deadline
.is_none_or(|deadline| deadline > snapshot_time)
});
let messages = normalize_snapshot_messages(&session.id, &mut cache, messages);
Some(ChatSnapshot {
state: cache.state.clone(),
transcript: ChatTranscriptPage {
messages,
tools: cache.tools.values().cloned().collect(),
older_cursor: None,
newer_cursor: None,
page_bytes: 0,
event_cursor: barrier,
generation: uuid::Uuid::new_v4().to_string(),
},
extension: ChatExtensionState {
pending_requests: cache
.pending_extension_requests
.iter()
.map(|pending| {
let mut record = pending.record.clone();
if let (Some(deadline), Some(payload)) =
(pending.deadline, record.payload.as_object_mut())
{
let milliseconds = u64::try_from(
deadline
.saturating_duration_since(snapshot_time)
.as_millis(),
)
.unwrap_or(u64::MAX)
.max(1);
payload.insert("timeout".into(), Value::from(milliseconds));
}
record
})
.collect(),
notifications: cache.extension_notifications.iter().cloned().collect(),
statuses: cache.extension_statuses.iter().cloned().collect(),
widgets: cache.extension_widgets.iter().cloned().collect(),
editor_text: cache.extension_editor_text.clone(),
},
event_sequence: barrier,
})
}
async fn build_history_snapshot(
&self,
session: &PiChatSession,
mut page: ParsedHistoryPage,
barrier: u64,
owner: &AuthenticatedRequestOwner,
owner_epoch: u64,
request_id: &str,
) -> AgentResult<Option<ChatSnapshot>> {
session.ensure_owner_live_for(owner).await?;
let lease = session.owner.lock().await;
if !owner.is_live() {
drop(lease);
session.ensure_owner_live_for(owner).await?;
return Err(invalid_transcript_cursor());
}
if lease.owner.as_ref() != Some(owner) || lease.epoch != owner_epoch {
return Ok(None);
}
let mut cache = session.cache.lock().await;
if cache.last_observed_event_sequence > barrier {
return Ok(None);
}
let snapshot_time = Instant::now();
cache.pending_extension_requests.retain(|pending| {
pending
.deadline
.is_none_or(|deadline| deadline > snapshot_time)
});
let live_messages = reconcile_persisted_live_messages(&mut cache, &mut page.messages);
let state = cache.state.clone();
let extension = snapshot_extension_state(&cache, snapshot_time);
let cache_tools = cache.tools.values().cloned().collect::<Vec<_>>();
let trim_input = HistorySnapshotTrimInput {
page: &page,
live_messages: &live_messages,
cache_tools: &cache_tools,
state: &state,
extension: &extension,
event_cursor: barrier,
session_id: &session.id,
request_id,
};
let trimmed = select_history_snapshot_trim(&trim_input)?;
page.trim_oldest_messages(trimmed);
cache.transcript_generation = uuid::Uuid::new_v4().to_string();
cache.transcript_cursors = TranscriptCursorRegistry::default();
let transcript = transcript_from_history_parts(
&mut cache,
&session.id,
owner,
&page,
live_messages,
barrier,
);
Ok(Some(finalize_snapshot_page_bytes(
session.id.clone(),
request_id.to_owned(),
ChatSnapshot {
state: cache.state.clone(),
transcript,
extension,
event_sequence: barrier,
},
)?))
}
fn remove_entry(&self, session_id: &str) -> Option<Arc<ManagedSession>> {
lock_std(&self.registry).sessions.remove(session_id)
}
fn activate_entry(
&self,
entry: &Arc<ManagedSession>,
on_activated: impl FnOnce(),
) -> AgentResult<()> {
let registry = lock_std(&self.registry);
let registered = registry
.sessions
.get(&entry.id)
.is_some_and(|current| Arc::ptr_eq(current, entry));
if !registered || self.shutting_down.load(Ordering::Acquire) {
return Err(initialization_cancelled());
}
let lifecycle = lock_std(&entry.lifecycle);
if lifecycle.closed {
return Err(initialization_cancelled());
}
entry.active.store(true, Ordering::Release);
on_activated();
Ok(())
}
fn publish_for_active_session(
&self,
session_id: &str,
session: &Arc<PiChatSession>,
event: SequencedChatEvent,
) -> bool {
self.publish_batch_for_active_session(session_id, session, vec![event])
}
fn publish_batch_for_active_session(
&self,
session_id: &str,
session: &Arc<PiChatSession>,
events: Vec<SequencedChatEvent>,
) -> bool {
let registry = lock_std(&self.registry);
let Some(entry) = registry.sessions.get(session_id) else {
return false;
};
if !entry.active.load(Ordering::Acquire)
|| entry
.session()
.is_none_or(|registered| !Arc::ptr_eq(®istered, session))
{
return false;
}
session.open_publications(&self.events, events)
}
fn ensure_active_session(
&self,
session_id: &str,
expected_entry: &Arc<ManagedSession>,
expected_session: &Arc<PiChatSession>,
) -> AgentResult<()> {
let registry = lock_std(&self.registry);
let Some(entry) = registry.sessions.get(session_id) else {
return Err(session_missing());
};
if !Arc::ptr_eq(entry, expected_entry)
|| !entry.active.load(Ordering::Acquire)
|| entry
.session()
.is_none_or(|registered| !Arc::ptr_eq(®istered, expected_session))
{
return Err(session_missing());
}
Ok(())
}
async fn ensure_active_owned_session(
&self,
session_id: &str,
entry: &Arc<ManagedSession>,
session: &Arc<PiChatSession>,
owner: &AuthenticatedRequestOwner,
) -> AgentResult<()> {
self.ensure_active_session(session_id, entry, session)?;
session.ensure_owned_by(owner).await?;
self.ensure_active_session(session_id, entry, session)
}
}
impl ManagedSession {
fn set_history_path(&self, path: Option<PathBuf>) {
*lock_std(&self.history_path) = path;
}
fn history_path(&self) -> Option<PathBuf> {
lock_std(&self.history_path).clone()
}
fn attach(&self, session: Arc<PiChatSession>) -> AgentResult<()> {
let mut lifecycle = lock_std(&self.lifecycle);
if lifecycle.closed || *self.cancel_tx.borrow() {
return Err(initialization_cancelled());
}
if lifecycle.session.is_some() {
return Err(session_exists());
}
lifecycle.session = Some(session);
Ok(())
}
fn session(&self) -> Option<Arc<PiChatSession>> {
lock_std(&self.lifecycle).session.clone()
}
fn install_pump(&self, start: impl FnOnce() -> JoinHandle<()>) -> AgentResult<()> {
let mut lifecycle = lock_std(&self.lifecycle);
if lifecycle.closed || *self.cancel_tx.borrow() {
return Err(initialization_cancelled());
}
lifecycle.pump = Some(start());
Ok(())
}
fn mark_closed(&self) {
let session = {
let mut lifecycle = lock_std(&self.lifecycle);
lifecycle.closed = true;
lifecycle.session.clone()
};
if let Some(session) = session {
session.closing.store(true, Ordering::Release);
session.discard_publications();
session.processed_event_notify.notify_waiters();
session.prompt_lifecycle_notify.notify_waiters();
}
let _ = self.cancel_tx.send(true);
}
fn take_resources(&self) -> (Option<Arc<PiChatSession>>, Option<JoinHandle<()>>) {
let mut lifecycle = lock_std(&self.lifecycle);
(lifecycle.session.take(), lifecycle.pump.take())
}
fn ensure_open(&self) -> AgentResult<()> {
if lock_std(&self.lifecycle).closed || *self.cancel_tx.borrow() {
Err(initialization_cancelled())
} else {
Ok(())
}
}
}
impl InitializationGuard {
fn new(inner: &Arc<ManagerInner>, entry: Arc<ManagedSession>) -> Self {
Self {
inner: Arc::downgrade(inner),
entry,
armed: true,
}
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for InitializationGuard {
fn drop(&mut self) {
if !self.armed {
return;
}
if let Some(inner) = self.inner.upgrade() {
inner.remove_matching(&self.entry.id, self.entry.generation);
}
let entry = self.entry.clone();
entry.mark_closed();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move { cleanup_managed(entry).await });
}
}
}
impl ActivePromptGuard {
fn new(inner: &Arc<ManagerInner>, entry: Arc<ManagedSession>) -> Self {
Self {
inner: Arc::downgrade(inner),
entry,
armed: true,
}
}
fn disarm(&mut self) {
self.armed = false;
}
async fn cleanup(&mut self) {
if !self.armed {
return;
}
if let Some(inner) = self.inner.upgrade() {
inner.remove_matching(&self.entry.id, self.entry.generation);
}
self.entry.mark_closed();
cleanup_managed(self.entry.clone()).await;
self.armed = false;
}
}
impl Drop for ActivePromptGuard {
fn drop(&mut self) {
if !self.armed {
return;
}
if let Some(inner) = self.inner.upgrade() {
inner.remove_matching(&self.entry.id, self.entry.generation);
}
let entry = self.entry.clone();
entry.mark_closed();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move { cleanup_managed(entry).await });
}
}
}
impl ManagerInner {
fn remove_matching(&self, session_id: &str, generation: u64) -> Option<Arc<ManagedSession>> {
let mut registry = lock_std(&self.registry);
if registry
.sessions
.get(session_id)
.is_some_and(|entry| entry.generation == generation)
{
registry.sessions.remove(session_id)
} else {
None
}
}
#[cfg(test)]
async fn pause_after_normalize_if_enabled(&self) {
if !self.normalize_pause_enabled.load(Ordering::Acquire)
|| self.normalize_pause_claimed.swap(true, Ordering::AcqRel)
{
return;
}
self.normalize_paused.store(true, Ordering::Release);
self.normalize_paused_notify.notify_waiters();
loop {
let notified = self.normalize_resume_notify.notified();
if !self.normalize_pause_enabled.load(Ordering::Acquire) {
return;
}
notified.await;
}
}
}
fn transcript_page_from_history(
cache: &mut SessionCache,
session_id: &str,
owner: &AuthenticatedRequestOwner,
page: ParsedHistoryPage,
event_cursor: u64,
) -> ChatTranscriptPage {
let live_messages = Vec::new();
transcript_from_history_parts(cache, session_id, owner, &page, live_messages, event_cursor)
}
fn transcript_from_history_parts(
cache: &mut SessionCache,
session_id: &str,
owner: &AuthenticatedRequestOwner,
page: &ParsedHistoryPage,
live_messages: Vec<ChatMessage>,
event_cursor: u64,
) -> ChatTranscriptPage {
let mut messages = page.messages.clone();
messages.extend(live_messages);
let mut tools = HashMap::new();
for tool in &page.tools {
tools.insert(tool.id.clone(), tool.clone());
}
for tool in cache.tools.values() {
tools.insert(tool.id.clone(), tool.clone());
}
let older_cursor = page
.has_older
.then(|| mint_transcript_cursor(cache, session_id, owner, &page.boundary));
let newer_cursor = page
.has_newer
.then(|| mint_transcript_cursor(cache, session_id, owner, &page.boundary));
ChatTranscriptPage {
messages,
tools: tools.into_values().collect(),
older_cursor,
newer_cursor,
page_bytes: 0,
event_cursor,
generation: cache.transcript_generation.clone(),
}
}
fn mint_transcript_cursor(
cache: &mut SessionCache,
session_id: &str,
owner: &AuthenticatedRequestOwner,
boundary: &HistoryPageBoundary,
) -> String {
while cache.transcript_cursors.records.len() >= MAX_TRANSCRIPT_CURSORS_PER_SESSION {
let Some(evicted) = cache.transcript_cursors.insertion_order.pop_front() else {
break;
};
cache.transcript_cursors.records.remove(&evicted);
}
loop {
let cursor = uuid::Uuid::new_v4().to_string();
if cache.transcript_cursors.records.contains_key(&cursor) {
continue;
}
cache
.transcript_cursors
.insertion_order
.push_back(cursor.clone());
cache.transcript_cursors.records.insert(
cursor.clone(),
TranscriptCursorRecord {
session_id: session_id.to_owned(),
owner: owner.clone(),
generation: cache.transcript_generation.clone(),
boundary: boundary.clone(),
},
);
return cursor;
}
}
fn reconcile_persisted_live_messages(
cache: &mut SessionCache,
persisted: &mut [ChatMessage],
) -> Vec<ChatMessage> {
for stable in persisted.iter() {
if let Some(provisional) = cache
.recent_live_messages
.iter()
.find(|live| messages_reconcile(live, stable))
.map(|live| live.id.clone())
{
remember_persisted_alias(cache, provisional, stable.id.clone());
}
}
let persisted_ids = persisted
.iter()
.map(|message| message.id.clone())
.collect::<HashSet<_>>();
cache
.recent_live_messages
.iter()
.filter_map(|live| {
let stable = cache.persisted_aliases.get(&live.id).map(String::as_str);
if stable.is_some_and(|id| persisted_ids.contains(id))
|| persisted.iter().any(|message| message.id == live.id)
{
return None;
}
Some(live.clone())
})
.collect()
}
fn messages_reconcile(provisional: &ChatMessage, persisted: &ChatMessage) -> bool {
if provisional.role != persisted.role {
return false;
}
if let (Some(left), Some(right)) = (provisional.timestamp, persisted.timestamp)
&& left != right
{
return false;
}
provisional.content.len() == persisted.content.len()
&& provisional
.content
.iter()
.zip(&persisted.content)
.all(|(left, right)| chat_content_reconciles(left, right))
}
fn chat_content_reconciles(left: &ChatContent, right: &ChatContent) -> bool {
match (left, right) {
(ChatContent::Text { text: left }, ChatContent::Text { text: right })
| (ChatContent::Thinking { text: left }, ChatContent::Thinking { text: right }) => {
left == right || left.starts_with(right) || right.starts_with(left)
}
_ => left == right,
}
}
fn remember_persisted_alias(cache: &mut SessionCache, provisional: String, stable: String) {
if cache.persisted_aliases.contains_key(&provisional) {
cache.persisted_alias_order.retain(|id| id != &provisional);
} else if cache.persisted_aliases.len() >= MAX_PERSISTED_ALIASES
&& let Some(evicted) = cache.persisted_alias_order.pop_front()
{
cache.persisted_aliases.remove(&evicted);
}
cache.persisted_alias_order.push_back(provisional.clone());
cache.persisted_aliases.insert(provisional, stable);
}
fn record_recent_live_message(cache: &mut SessionCache, message: ChatMessage) {
if let Some(existing) = cache
.recent_live_messages
.iter_mut()
.find(|existing| existing.id == message.id)
{
*existing = message;
return;
}
push_bounded(
&mut cache.recent_live_messages,
message,
MAX_RECENT_LIVE_MESSAGES,
);
}
fn snapshot_extension_state(cache: &SessionCache, snapshot_time: Instant) -> ChatExtensionState {
ChatExtensionState {
pending_requests: cache
.pending_extension_requests
.iter()
.map(|pending| {
let mut record = pending.record.clone();
if let (Some(deadline), Some(payload)) =
(pending.deadline, record.payload.as_object_mut())
{
let milliseconds = u64::try_from(
deadline
.saturating_duration_since(snapshot_time)
.as_millis(),
)
.unwrap_or(u64::MAX)
.max(1);
payload.insert("timeout".into(), Value::from(milliseconds));
}
record
})
.collect(),
notifications: cache.extension_notifications.iter().cloned().collect(),
statuses: cache.extension_statuses.iter().cloned().collect(),
widgets: cache.extension_widgets.iter().cloned().collect(),
editor_text: cache.extension_editor_text.clone(),
}
}
const SNAPSHOT_OPAQUE_TOKEN_BYTES: usize = 36;
struct HistorySnapshotTrimInput<'a> {
page: &'a ParsedHistoryPage,
live_messages: &'a [ChatMessage],
cache_tools: &'a [ChatToolExecution],
state: &'a ChatSessionState,
extension: &'a ChatExtensionState,
event_cursor: u64,
session_id: &'a str,
request_id: &'a str,
}
fn select_history_snapshot_trim(input: &HistorySnapshotTrimInput<'_>) -> AgentResult<usize> {
let persisted_count = input.page.messages.len();
if history_snapshot_fits(input, 0)? {
return Ok(0);
}
if !history_snapshot_fits(input, persisted_count)? {
return Err(transcript_snapshot_too_large());
}
let mut low = 1_usize;
let mut high = persisted_count;
while low < high {
let mid = low + (high - low) / 2;
if history_snapshot_fits(input, mid)? {
high = mid;
} else {
low = mid.saturating_add(1);
}
}
Ok(low)
}
fn history_snapshot_fits(
input: &HistorySnapshotTrimInput<'_>,
trim_count: usize,
) -> AgentResult<bool> {
let mut page = input.page.clone();
page.trim_oldest_messages(trim_count);
let transcript = preview_history_transcript(
&page,
input.live_messages,
input.cache_tools,
input.event_cursor,
);
let snapshot = ChatSnapshot {
state: input.state.clone(),
transcript,
extension: input.extension.clone(),
event_sequence: input.event_cursor,
};
match finalize_snapshot_page_bytes(
input.session_id.to_owned(),
input.request_id.to_owned(),
snapshot,
) {
Ok(_) => Ok(true),
Err(error) if error.code() == ErrorCode::TranscriptSnapshotTooLarge => Ok(false),
Err(error) => Err(error),
}
}
fn preview_history_transcript(
page: &ParsedHistoryPage,
live_messages: &[ChatMessage],
cache_tools: &[ChatToolExecution],
event_cursor: u64,
) -> ChatTranscriptPage {
let mut messages = page.messages.clone();
messages.extend(live_messages.iter().cloned());
let mut tools = HashMap::new();
for tool in &page.tools {
tools.insert(tool.id.clone(), tool.clone());
}
for tool in cache_tools {
tools.insert(tool.id.clone(), tool.clone());
}
ChatTranscriptPage {
messages,
tools: tools.into_values().collect(),
older_cursor: page
.has_older
.then(|| "x".repeat(SNAPSHOT_OPAQUE_TOKEN_BYTES)),
newer_cursor: page
.has_newer
.then(|| "x".repeat(SNAPSHOT_OPAQUE_TOKEN_BYTES)),
page_bytes: 0,
event_cursor,
generation: "x".repeat(SNAPSHOT_OPAQUE_TOKEN_BYTES),
}
}
fn finalize_snapshot_page_bytes(
session_id: String,
request_id: String,
mut snapshot: ChatSnapshot,
) -> AgentResult<ChatSnapshot> {
for _ in 0..16 {
retain_referenced_tools(&mut snapshot.transcript);
let message =
AgentMessage::chat_snapshot(session_id.clone(), request_id.clone(), snapshot.clone());
let bytes = serde_json::to_vec(&message)
.map(|serialized| serialized.len())
.map_err(|_| transcript_snapshot_too_large())?;
if bytes > MAX_TRANSCRIPT_DELIVERY_BYTES {
return Err(transcript_snapshot_too_large());
}
if snapshot.transcript.page_bytes == bytes {
return Ok(snapshot);
}
snapshot.transcript.page_bytes = bytes;
}
Err(transcript_snapshot_too_large())
}
fn finalize_chat_snapshot(
session_id: String,
request_id: String,
snapshot: ChatSnapshot,
) -> AgentResult<AgentMessage> {
let snapshot = finalize_snapshot_page_bytes(session_id.clone(), request_id.clone(), snapshot)?;
Ok(AgentMessage::chat_snapshot(
session_id, request_id, snapshot,
))
}
fn retain_referenced_tools(page: &mut ChatTranscriptPage) {
let referenced = page
.messages
.iter()
.flat_map(|message| message.content.iter())
.filter_map(|content| match content {
ChatContent::ToolCall { id, .. } => Some(id.as_str()),
ChatContent::ToolResult { tool_call_id, .. } => Some(tool_call_id.as_str()),
_ => None,
})
.collect::<HashSet<_>>();
page.tools
.retain(|tool| referenced.contains(tool.id.as_str()));
}
fn finalize_transcript_page(
session_id: String,
request_id: String,
mut page: ChatTranscriptPage,
) -> AgentResult<AgentMessage> {
for _ in 0..16 {
let message = AgentMessage::ChatTranscriptPage {
session_id: session_id.clone(),
request_id: request_id.clone(),
page: page.clone(),
};
let bytes = serde_json::to_vec(&message)
.map(|serialized| serialized.len())
.map_err(|_| transcript_page_too_large())?;
if bytes > MAX_TRANSCRIPT_DELIVERY_BYTES {
return Err(transcript_page_too_large());
}
if page.page_bytes == bytes {
return Ok(message);
}
page.page_bytes = bytes;
}
Err(AgentError::new(
ErrorCode::HistoryReadFailed,
"failed to stabilize transcript page envelope size",
))
}
fn map_transcript_page_error(error: AgentError) -> AgentError {
if error.code() == ErrorCode::HistoryReadFailed
&& error.message().starts_with("Pi transcript page is stale")
{
stale_transcript_cursor()
} else {
error
}
}
fn history_task_failed() -> AgentError {
AgentError::new(
ErrorCode::BackendDisconnected,
"transcript history page task stopped",
)
}
fn invalid_transcript_cursor() -> AgentError {
AgentError::new(ErrorCode::InvalidMessage, "invalid transcript cursor")
}
fn stale_prompt_cancellation() -> AgentError {
AgentError::new(
ErrorCode::InvalidMessage,
"prompt cancellation target was not observed before timeout",
)
}
#[cfg(test)]
fn internal_request_owner() -> AuthenticatedRequestOwner {
AuthenticatedRequestOwner::legacy("internal-manager-caller".into())
}
fn stale_transcript_cursor() -> AgentError {
AgentError::new(
ErrorCode::StaleTranscriptCursor,
"reload the newest transcript page",
)
}
fn transcript_page_too_large() -> AgentError {
AgentError::new(
ErrorCode::HistoryReadFailed,
"Pi transcript record exceeds the supported page record size",
)
}
fn transcript_snapshot_too_large() -> AgentError {
AgentError::new(
ErrorCode::TranscriptSnapshotTooLarge,
"transcript snapshot is too large; request a fresh transcript resync",
)
}
async fn cleanup_managed(entry: Arc<ManagedSession>) {
entry.mark_closed();
let (session, pump) = entry.take_resources();
if let Some(session) = session {
session.clear_prompt_lifecycle_on_close().await;
let _ = session.process.shutdown().await;
}
if let Some(mut pump) = pump
&& timeout_at(Instant::now() + Duration::from_secs(1), &mut pump)
.await
.is_err()
{
pump.abort();
let _ = pump.await;
}
}
fn lock_std<T>(mutex: &StdMutex<T>) -> std::sync::MutexGuard<'_, T> {
mutex.lock().unwrap_or_else(|error| error.into_inner())
}
fn normalized_history_path(workspace: &Path, path: &str) -> PathBuf {
let path = PathBuf::from(path);
let path = if path.is_absolute() {
path
} else {
workspace.join(path)
};
path.canonicalize().unwrap_or(path)
}
impl PiChatSession {
fn starting(
id: String,
workspace: PathBuf,
process: Arc<PiRpcProcess>,
history: Option<PersistedTranscript>,
owner: AuthenticatedRequestOwner,
) -> Self {
let state = ChatSessionState {
session_id: id.clone(),
pi_session_id: String::new(),
session_file: None,
workspace: workspace.clone(),
model: None,
thinking_level: ThinkingLevel::Off,
status: ChatStatus::Starting,
is_streaming: false,
is_compacting: false,
name: None,
message_count: 0,
pending_message_count: 0,
stats: None,
};
Self {
id,
workspace,
process,
cache: Mutex::new(SessionCache {
state,
prompt_pending: false,
active_prompt: None,
cancellation_intents: HashSet::new(),
cancellation_intent_order: VecDeque::new(),
canceled_message_ids: VecDeque::new(),
next_message_id: 1,
active_message_id: None,
last_assistant_message: None,
pi_message_ids: HashMap::new(),
snapshot_message_ids: HashMap::new(),
tools: HashMap::new(),
terminal_error_emitted: false,
last_event_sequence: 0,
last_observed_event_sequence: 0,
live_messages: HashMap::new(),
recent_live_messages: VecDeque::new(),
persisted_aliases: HashMap::new(),
persisted_alias_order: VecDeque::new(),
transcript_generation: uuid::Uuid::new_v4().to_string(),
transcript_cursors: TranscriptCursorRegistry::default(),
pending_user_messages: VecDeque::new(),
terminal_status: None,
pending_extension_requests: VecDeque::new(),
extension_notifications: VecDeque::new(),
extension_statuses: VecDeque::new(),
extension_widgets: VecDeque::new(),
extension_editor_text: None,
extension_title: None,
}),
closing: AtomicBool::new(false),
processed_event_sequence: AtomicU64::new(0),
processed_event_notify: Notify::new(),
prompt_lifecycle_notify: Notify::new(),
publication_gate: Arc::new(Mutex::new(())),
publication_state: StdMutex::new(PublicationState::Gated {
events: VecDeque::new(),
bytes: 0,
overflowed: false,
}),
history,
owner: Mutex::new(OwnerLease {
owner: Some(owner),
epoch: 1,
}),
}
}
#[cfg(test)]
async fn owner(&self) -> AuthenticatedRequestOwner {
self.owner
.lock()
.await
.owner
.clone()
.expect("test session has an active transcript owner")
}
async fn ensure_owner_live(&self) -> AgentResult<()> {
let owner = self
.owner
.lock()
.await
.owner
.clone()
.ok_or_else(session_missing)?;
self.ensure_owner_live_for(&owner).await
}
async fn ensure_owner_live_for(&self, owner: &AuthenticatedRequestOwner) -> AgentResult<()> {
if owner.is_live() {
return Ok(());
}
self.release_transcript_owner(owner).await;
owner.ensure_live()
}
async fn ensure_owned_by(&self, owner: &AuthenticatedRequestOwner) -> AgentResult<()> {
owner.ensure_live()?;
let lease = self.owner.lock().await;
if lease.owner.as_ref() == Some(owner) {
Ok(())
} else {
Err(session_missing())
}
}
async fn owner_epoch(&self, owner: &AuthenticatedRequestOwner) -> AgentResult<u64> {
let lease = self.owner.lock().await;
if lease.owner.as_ref() == Some(owner) {
Ok(lease.epoch)
} else {
Err(invalid_transcript_cursor())
}
}
async fn claim_snapshot_owner(&self, owner: &AuthenticatedRequestOwner) -> AgentResult<u64> {
let (epoch, owner_changed) = {
let mut lease = self.owner.lock().await;
owner.ensure_live()?;
match &lease.owner {
Some(current) if current == owner => (lease.epoch, false),
Some(current) if current.can_replace_after_ui_reconnect(owner) => {
current.invalidate();
lease.owner = Some(owner.clone());
lease.epoch = lease.epoch.saturating_add(1).max(1);
(lease.epoch, true)
}
Some(current) if current.can_transfer_to_other_ui(owner) => {
lease.owner = Some(owner.clone());
lease.epoch = lease.epoch.saturating_add(1).max(1);
(lease.epoch, true)
}
Some(_) => return Err(invalid_transcript_cursor()),
None => {
lease.owner = Some(owner.clone());
lease.epoch = lease.epoch.saturating_add(1).max(1);
(lease.epoch, true)
}
}
};
if owner_changed {
let mut cache = self.cache.lock().await;
cache.transcript_generation = uuid::Uuid::new_v4().to_string();
cache.transcript_cursors = TranscriptCursorRegistry::default();
}
Ok(epoch)
}
async fn release_transcript_owner(&self, owner: &AuthenticatedRequestOwner) {
let released = {
let mut lease = self.owner.lock().await;
if lease.owner.as_ref() != Some(owner) {
false
} else {
lease.owner = None;
lease.epoch = lease.epoch.saturating_add(1).max(1);
true
}
};
if released {
let mut cache = self.cache.lock().await;
cache.transcript_generation = uuid::Uuid::new_v4().to_string();
cache.transcript_cursors = TranscriptCursorRegistry::default();
}
}
fn begin_publications(&self) -> AgentResult<()> {
let mut state = lock_std(&self.publication_state);
match &*state {
PublicationState::Open => {
*state = PublicationState::Gated {
events: VecDeque::new(),
bytes: 0,
overflowed: false,
};
Ok(())
}
PublicationState::Gated { .. } => Err(AgentError::new(
ErrorCode::BackendDisconnected,
"Pi chat publication batch is already active",
)),
PublicationState::Discarded => Err(session_missing()),
}
}
fn publish_events(
&self,
sender: &broadcast::Sender<SequencedChatEvent>,
events: Vec<SequencedChatEvent>,
) {
let mut state = lock_std(&self.publication_state);
match &mut *state {
PublicationState::Open => publish(sender, events),
PublicationState::Gated {
events: buffered,
bytes,
overflowed,
} => {
for event in events {
if *overflowed {
continue;
}
let (event, event_bytes) = bounded_fanout_event(event);
if bytes.saturating_add(event_bytes) > MAX_GATED_PUBLICATION_BYTES {
buffered.clear();
let (event, event_bytes) = bounded_fanout_event(SequencedChatEvent {
event_sequence: event.event_sequence,
event: ChatEvent::ResyncRequired {
session_id: self.id.clone(),
reason: "gated chat events exceeded the 16 MiB buffer; request a snapshot"
.into(),
},
});
buffered.push_back(event);
*bytes = event_bytes;
*overflowed = true;
} else {
buffered.push_back(event);
*bytes = bytes.saturating_add(event_bytes);
}
}
}
PublicationState::Discarded => {}
}
}
fn open_publications(
&self,
sender: &broadcast::Sender<SequencedChatEvent>,
prefix: Vec<SequencedChatEvent>,
) -> bool {
let mut state = lock_std(&self.publication_state);
let prior = std::mem::replace(&mut *state, PublicationState::Open);
match prior {
PublicationState::Gated { events, .. } => {
publish(sender, prefix);
publish(sender, events.into_iter().collect());
true
}
PublicationState::Open => false,
PublicationState::Discarded => {
*state = PublicationState::Discarded;
false
}
}
}
fn discard_publications(&self) {
*lock_std(&self.publication_state) = PublicationState::Discarded;
}
async fn state(&self) -> ChatSessionState {
self.cache.lock().await.state.clone()
}
async fn ensure_prompt_is_idle(&self) -> AgentResult<()> {
let cache = self.cache.lock().await;
if cache.state.is_streaming
|| cache.state.status == ChatStatus::Streaming
|| cache.active_message_id.is_some()
|| cache.active_prompt.is_some()
{
return Err(AgentError::new(
ErrorCode::InvalidMessage,
"chat prompt cannot be submitted while Pi is streaming",
));
}
Ok(())
}
async fn update_from_pi_state(&self, state: &PiSessionState, barrier: u64) {
let mut cache = self.cache.lock().await;
let previous_status = cache.state.status;
let active_prompt_status = cache.active_prompt.as_ref().map(|active| match active.phase {
PromptPhase::Preflight { .. } => ChatStatus::Starting,
PromptPhase::Running | PromptPhase::Canceling
if previous_status == ChatStatus::Retrying =>
{
ChatStatus::Retrying
}
PromptPhase::Running | PromptPhase::Canceling => ChatStatus::Streaming,
});
cache.state.pi_session_id = state.session_id.clone();
cache.state.session_file = state.session_file.clone();
cache.state.workspace = self.workspace.clone();
if cache.last_observed_event_sequence > barrier {
return;
}
cache.state.model = state
.model
.clone()
.and_then(|value| model_from_value(value, None));
cache.state.thinking_level = state.thinking_level;
cache.state.is_streaming = state.is_streaming;
cache.state.is_compacting = state.is_compacting;
cache.state.name = cache
.extension_title
.clone()
.or_else(|| state.session_name.clone());
cache.state.message_count = state.message_count;
cache.state.pending_message_count = state.pending_message_count;
let pi_status = status_from_state(state);
cache.state.status = cache.terminal_status.unwrap_or_else(|| {
if pi_status == ChatStatus::Idle {
active_prompt_status.unwrap_or(ChatStatus::Idle)
} else {
pi_status
}
});
if cache.state.status == ChatStatus::Streaming && active_prompt_status.is_some() {
cache.state.is_streaming = true;
}
}
async fn set_model(&self, model: Option<ChatModel>) {
self.cache.lock().await.state.model = model;
}
async fn set_thinking(&self, level: ThinkingLevel) {
self.cache.lock().await.state.thinking_level = level;
}
async fn mark_initialized(&self) {
let mut cache = self.cache.lock().await;
if cache.state.status == ChatStatus::Starting {
cache.state.status = ChatStatus::Idle;
}
}
async fn prepare_prompt(&self, request_id: &str, text: &str) -> PromptAccepted {
let mut cache = self.cache.lock().await;
let user_message_id = next_message_id(&mut cache, &self.id);
cache
.pending_user_messages
.push_back((text.to_owned(), user_message_id.clone()));
let message = ChatMessage {
id: user_message_id.clone(),
role: ChatRole::User,
content: vec![ChatContent::Text {
text: text.to_owned(),
}],
timestamp: None,
provider: None,
model: None,
usage: None,
stop_reason: None,
error_message: None,
};
cache
.live_messages
.insert(user_message_id.clone(), message.clone());
record_recent_live_message(&mut cache, message);
let canceled_before_start = take_cancellation_intent(&mut cache, request_id);
cache.active_prompt = Some(ActivePromptLifecycle {
request_id: request_id.to_owned(),
user_message_id: user_message_id.clone(),
kind: PromptKind::Unknown,
phase: PromptPhase::Preflight { dispatched: false },
canceled_before_start,
abort_sent: false,
});
cache.terminal_status = None;
sync_prompt_pending(&mut cache);
let accepted = PromptAccepted {
request_id: request_id.to_owned(),
user_message_id,
};
drop(cache);
self.prompt_lifecycle_notify.notify_waiters();
accepted
}
async fn set_prompt_kind(&self, request_id: &str, kind: PromptKind) {
let mut cache = self.cache.lock().await;
if let Some(active) = cache
.active_prompt
.as_mut()
.filter(|active| active.request_id == request_id)
{
active.kind = kind;
}
}
async fn reject_prompt(&self, accepted: &PromptAccepted) {
let mut cache = self.cache.lock().await;
clear_prompt_lifecycle(&mut cache, &accepted.request_id);
cache
.pending_user_messages
.retain(|(_, id)| id != &accepted.user_message_id);
cache.live_messages.remove(&accepted.user_message_id);
cache
.recent_live_messages
.retain(|message| message.id != accepted.user_message_id);
drop(cache);
self.prompt_lifecycle_notify.notify_waiters();
}
async fn begin_prompt_dispatch(&self, request_id: &str) -> PromptDispatch {
let dispatch = {
let mut cache = self.cache.lock().await;
let active = cache
.active_prompt
.as_mut()
.filter(|active| active.request_id == request_id);
match active {
Some(active) => match active.phase {
PromptPhase::Preflight { dispatched: false }
if active.canceled_before_start =>
{
PromptDispatch::CanceledBeforeStart
}
PromptPhase::Preflight { dispatched: false } => {
active.phase = PromptPhase::Preflight { dispatched: true };
PromptDispatch::SendToPi
}
_ => PromptDispatch::CanceledBeforeStart,
},
None => PromptDispatch::CanceledBeforeStart,
}
};
self.prompt_lifecycle_notify.notify_waiters();
dispatch
}
async fn mark_prompt_running(&self, request_id: &str) {
let mut cache = self.cache.lock().await;
if let Some(active) = cache
.active_prompt
.as_mut()
.filter(|active| active.request_id == request_id)
&& matches!(active.phase, PromptPhase::Preflight { .. })
{
active.phase = PromptPhase::Running;
}
drop(cache);
self.prompt_lifecycle_notify.notify_waiters();
}
async fn mark_prompt_started(&self) {
let mut cache = self.cache.lock().await;
cache.terminal_status = None;
if let Some(active) = cache.active_prompt.as_mut()
&& matches!(active.phase, PromptPhase::Preflight { .. })
{
active.phase = if active.canceled_before_start {
PromptPhase::Canceling
} else {
PromptPhase::Running
};
}
drop(cache);
self.prompt_lifecycle_notify.notify_waiters();
}
async fn complete_preflight_cancellation(&self, accepted: &PromptAccepted) -> bool {
let completed = {
let mut cache = self.cache.lock().await;
let matches_lifecycle = cache.active_prompt.as_ref().is_some_and(|active| {
active.request_id == accepted.request_id
&& active.user_message_id == accepted.user_message_id
&& active.canceled_before_start
&& matches!(active.phase, PromptPhase::Preflight { dispatched: false })
});
if matches_lifecycle {
mark_canceled_message(&mut cache, &accepted.user_message_id);
cache.state.status = ChatStatus::Settled;
cache.state.is_streaming = false;
cache.state.is_compacting = false;
cache.terminal_status = Some(ChatStatus::Settled);
clear_prompt_lifecycle(&mut cache, &accepted.request_id);
}
matches_lifecycle
};
self.prompt_lifecycle_notify.notify_waiters();
completed
}
async fn complete_synchronous_extension(&self, accepted: &PromptAccepted) -> bool {
let completed = {
let mut cache = self.cache.lock().await;
let matches_lifecycle = cache.active_prompt.as_ref().is_some_and(|active| {
active.request_id == accepted.request_id
&& active.user_message_id == accepted.user_message_id
&& active.kind == PromptKind::SynchronousExtension
&& matches!(active.phase, PromptPhase::Preflight { .. })
});
if matches_lifecycle {
clear_prompt_lifecycle(&mut cache, &accepted.request_id);
cache.state.status = ChatStatus::Settled;
cache.state.is_streaming = false;
cache.state.is_compacting = false;
cache.terminal_status = None;
}
matches_lifecycle
};
if completed {
self.prompt_lifecycle_notify.notify_waiters();
}
completed
}
async fn prompt_lifecycle(&self, request_id: &str) -> bool {
self.cache
.lock()
.await
.active_prompt
.as_ref()
.is_some_and(|active| active.request_id == request_id)
}
async fn has_cancellation_intent(&self, request_id: &str) -> bool {
self.cache
.lock()
.await
.cancellation_intents
.contains(request_id)
}
async fn remove_cancellation_intent(&self, request_id: &str) {
let mut cache = self.cache.lock().await;
remove_cancellation_intent(&mut cache, request_id);
drop(cache);
self.prompt_lifecycle_notify.notify_waiters();
}
async fn expire_cancellation_intent_if_lifecycle_absent(&self, request_id: &str) -> bool {
let mut cache = self.cache.lock().await;
let lifecycle_present = cache
.active_prompt
.as_ref()
.is_some_and(|active| active.request_id == request_id);
if !lifecycle_present {
remove_cancellation_intent(&mut cache, request_id);
}
drop(cache);
if !lifecycle_present {
self.prompt_lifecycle_notify.notify_waiters();
}
lifecycle_present
}
async fn abort_wait_failed(&self) -> bool {
self.closing.load(Ordering::Acquire) || self.cache.lock().await.terminal_error_emitted
}
async fn cancel_prompt_lifecycle(&self, request_id: &str) -> PromptCancellation {
let (cancellation, changed) = {
let mut cache = self.cache.lock().await;
if let Some(active) = cache
.active_prompt
.as_mut()
.filter(|active| active.request_id == request_id)
{
match active.phase {
PromptPhase::Preflight { .. } => {
let changed = !active.canceled_before_start;
active.canceled_before_start = true;
(PromptCancellation::Wait, changed)
}
PromptPhase::Running if active.abort_sent => (PromptCancellation::Wait, false),
PromptPhase::Running => {
active.abort_sent = true;
active.phase = PromptPhase::Canceling;
(PromptCancellation::SendAbortToPi, true)
}
PromptPhase::Canceling if active.abort_sent => {
(PromptCancellation::Wait, false)
}
PromptPhase::Canceling => {
active.abort_sent = true;
(PromptCancellation::SendAbortToPi, true)
}
}
} else {
let changed = record_cancellation_intent(&mut cache, request_id);
(PromptCancellation::Wait, changed)
}
};
if changed {
self.prompt_lifecycle_notify.notify_waiters();
}
cancellation
}
async fn wait_for_event_barrier(&self, barrier: u64, wait: Duration) -> AgentResult<()> {
let deadline = Instant::now() + wait;
loop {
let notified = self.processed_event_notify.notified();
if self.processed_event_sequence.load(Ordering::Acquire) >= barrier {
return Ok(());
}
if self.closing.load(Ordering::Acquire) {
return Err(AgentError::new(
ErrorCode::BackendDisconnected,
"Pi chat event pump closed before the response barrier",
));
}
if timeout_at(deadline, notified).await.is_err() {
return Err(AgentError::new(
ErrorCode::BackendDisconnected,
format!("Pi chat event barrier {barrier} timed out"),
));
}
}
}
async fn mark_event_processed(&self, sequence: u64) {
self.cache.lock().await.last_event_sequence = sequence;
self.processed_event_sequence
.fetch_max(sequence, Ordering::AcqRel);
self.processed_event_notify.notify_waiters();
}
async fn clear_prompt_lifecycle_on_close(&self) {
let mut cache = self.cache.lock().await;
clear_active_prompt_lifecycle(&mut cache);
drop(cache);
self.prompt_lifecycle_notify.notify_waiters();
}
}
async fn run_event_pump(
context: EventPumpContext,
mut receiver: PiRpcEventReceiver,
mut cancel_rx: watch::Receiver<bool>,
) {
let EventPumpContext {
manager,
generation,
session,
sender,
request_timeout,
failure_sequence,
pump_delay,
} = context;
loop {
let sequenced = tokio::select! {
biased;
changed = cancel_rx.changed() => {
let _ = changed;
return;
}
event = receiver.recv_sequenced() => event,
};
let Some(sequenced) = sequenced else { break };
if !pump_delay.is_zero() {
tokio::select! {
biased;
changed = cancel_rx.changed() => {
let _ = changed;
return;
}
_ = tokio::time::sleep(pump_delay) => {}
}
}
let normalized = normalize_event(
&session,
sequenced.sequence,
sequenced.event,
request_timeout,
)
.await;
#[cfg(test)]
if let Some(manager) = manager.upgrade() {
manager.pause_after_normalize_if_enabled().await;
}
match normalized {
Ok(events) => {
if *cancel_rx.borrow() || session.closing.load(Ordering::Acquire) {
return;
}
session.publish_events(
&sender,
events
.into_iter()
.map(|event| SequencedChatEvent {
event_sequence: Some(sequenced.sequence),
event,
})
.collect(),
);
session.mark_event_processed(sequenced.sequence).await;
}
Err(error) => {
publish_terminal_error(&session, &sender, &error, false, Some(sequenced.sequence))
.await;
session.mark_event_processed(sequenced.sequence).await;
session.closing.store(true, Ordering::Release);
let _ = session.process.shutdown().await;
if let Some(manager) = manager.upgrade() {
manager.remove_matching(&session.id, generation);
}
return;
}
}
}
if session.closing.load(Ordering::Acquire) {
return;
}
let failure = session
.process
.request(
PiRpcCommand::GetState {
id: format!(
"__regy:{}:event-pump-failure:{failure_sequence}",
session.id
),
},
request_timeout,
)
.await
.err()
.unwrap_or_else(|| {
AgentError::new(
ErrorCode::BackendDisconnected,
"Pi RPC event receiver closed unexpectedly",
)
});
if *cancel_rx.borrow() {
return;
}
publish_terminal_error(&session, &sender, &failure, true, None).await;
session.closing.store(true, Ordering::Release);
session.processed_event_notify.notify_waiters();
if let Some(manager) = manager.upgrade() {
manager.remove_matching(&session.id, generation);
}
}
fn publish(sender: &broadcast::Sender<SequencedChatEvent>, events: Vec<SequencedChatEvent>) {
for event in events {
publish_event(sender, event);
}
}
fn publish_event(sender: &broadcast::Sender<SequencedChatEvent>, event: SequencedChatEvent) {
let (event, _) = bounded_fanout_event(event);
let _ = sender.send(event);
}
fn bounded_fanout_event(event: SequencedChatEvent) -> (SequencedChatEvent, usize) {
let size = serde_json::to_vec(&event).map(|bytes| bytes.len());
if let Ok(bytes) = size
&& bytes <= MAX_FANOUT_EVENT_BYTES
{
return (event, bytes);
}
let resync = SequencedChatEvent {
event_sequence: event.event_sequence,
event: ChatEvent::ResyncRequired {
session_id: event.event.session_id().to_owned(),
reason: "normalized chat event exceeded the 256 KiB fanout limit; request a snapshot"
.into(),
},
};
let bytes = serde_json::to_vec(&resync)
.map(|bytes| bytes.len())
.unwrap_or_default();
(resync, bytes)
}
async fn publish_terminal_error(
session: &PiChatSession,
sender: &broadcast::Sender<SequencedChatEvent>,
error: &AgentError,
exited: bool,
event_sequence: Option<u64>,
) {
let mut cache = session.cache.lock().await;
if cache.terminal_error_emitted {
return;
}
cache.terminal_error_emitted = true;
clear_active_prompt_lifecycle(&mut cache);
cache.state.status = if exited {
ChatStatus::Exited
} else {
ChatStatus::Error
};
cache.state.is_streaming = false;
cache.state.is_compacting = false;
drop(cache);
session.prompt_lifecycle_notify.notify_waiters();
let mut events = vec![
ChatEvent::StatusChanged {
session_id: session.id.clone(),
status: ChatStatus::Error,
},
ChatEvent::Error {
session_id: session.id.clone(),
code: error.code().as_str().into(),
message: error.message().into(),
recoverable: false,
},
];
if exited {
events.extend([
ChatEvent::StatusChanged {
session_id: session.id.clone(),
status: ChatStatus::Exited,
},
ChatEvent::Exited {
session_id: session.id.clone(),
message: error.message().into(),
},
]);
}
session.publish_events(
sender,
events
.into_iter()
.map(|event| SequencedChatEvent {
event_sequence,
event,
})
.collect(),
);
}
async fn normalize_event(
session: &PiChatSession,
sequence: u64,
event: PiRpcEvent,
request_timeout: Duration,
) -> AgentResult<Vec<ChatEvent>> {
{
let mut cache = session.cache.lock().await;
cache.last_observed_event_sequence = cache.last_observed_event_sequence.max(sequence);
}
let id = session.id.clone();
match event {
PiRpcEvent::AgentStart | PiRpcEvent::TurnStart => {
session.mark_prompt_started().await;
set_status(session, ChatStatus::Streaming, true, false).await;
Ok(vec![status_event(&id, ChatStatus::Streaming)])
}
PiRpcEvent::AgentEnd { will_retry, .. } => {
if !will_retry {
return Ok(vec![ChatEvent::AgentEnded {
session_id: id,
will_retry,
}]);
}
session.cache.lock().await.terminal_status = None;
set_status(session, ChatStatus::Retrying, false, false).await;
Ok(vec![
ChatEvent::AgentEnded {
session_id: id.clone(),
will_retry,
},
status_event(&id, ChatStatus::Retrying),
])
}
PiRpcEvent::AgentSettled => {
let (canceled_before_start, terminal_status) = {
let mut cache = session.cache.lock().await;
let terminal_status = cache.terminal_status;
let canceled = clear_active_prompt_lifecycle(&mut cache)
.and_then(|active| {
active
.canceled_before_start
.then_some((active.request_id, active.user_message_id))
})
.filter(|_| terminal_status.is_none_or(|status| status == ChatStatus::Settled));
if let Some((_, user_message_id)) = &canceled {
mark_canceled_message(&mut cache, user_message_id);
}
(canceled, terminal_status)
};
session.prompt_lifecycle_notify.notify_waiters();
let stats = session
.process
.request(
PiRpcCommand::GetSessionStats {
id: format!("__regy:{}:settled-stats:{sequence}", session.id),
},
request_timeout.min(MAX_STATS_REQUEST_TIMEOUT),
)
.await
.and_then(expect_stats)
.ok();
if let Some(stats) = &stats {
session.cache.lock().await.state.stats = Some(Box::new(stats.clone()));
}
let publish_deferred_settlement =
canceled_before_start.is_some() && terminal_status == Some(ChatStatus::Settled);
let mut events = Vec::with_capacity(3);
if let Some(stats) = stats {
events.push(ChatEvent::StatsUpdated {
session_id: id.clone(),
stats,
});
}
if let Some((request_id, user_message_id)) = canceled_before_start {
events.push(ChatEvent::PromptCanceled {
session_id: id.clone(),
request_id,
user_message_id,
});
}
if terminal_status.is_some() && !publish_deferred_settlement {
return Ok(events);
}
set_status(session, ChatStatus::Settled, false, false).await;
events.push(status_event(&id, ChatStatus::Settled));
Ok(events)
}
PiRpcEvent::TurnEnd { message, .. } => {
let cache = session.cache.lock().await;
let message_id = cache
.active_message_id
.clone()
.unwrap_or_else(|| format!("{}:turn", session.id));
Ok(usage_event(&id, &message_id, &message)
.into_iter()
.collect())
}
PiRpcEvent::MessageStart { message } => {
let mut cache = session.cache.lock().await;
let message_id = stream_message_id(&mut cache, &id, &message, true)?;
let normalized = normalize_message(message_id, message);
if normalized.role == ChatRole::Assistant {
cache.last_assistant_message = Some(normalized.clone());
}
cache
.live_messages
.insert(normalized.id.clone(), normalized.clone());
record_recent_live_message(&mut cache, normalized.clone());
Ok(vec![ChatEvent::MessageStarted {
session_id: id,
message: normalized,
}])
}
PiRpcEvent::MessageUpdate {
message,
assistant_message_event,
} => normalize_message_update(session, message, assistant_message_event).await,
PiRpcEvent::MessageEnd { message } => {
let identity = message_identity(&message);
let mut cache = session.cache.lock().await;
let message_id = stream_message_id(&mut cache, &id, &message, false)?;
register_stream_message_id(&mut cache, &identity, &message_id);
cache.active_message_id = None;
cache.state.message_count = cache.state.message_count.saturating_add(1);
let normalized = normalize_message(message_id.clone(), message);
if normalized.role == ChatRole::Assistant {
cache.last_assistant_message = Some(normalized.clone());
}
cache.live_messages.remove(&message_id);
record_recent_live_message(&mut cache, normalized.clone());
let mut events = Vec::new();
if let Some(usage) = normalized.usage.clone() {
events.push(ChatEvent::UsageUpdated {
session_id: id.clone(),
message_id,
usage,
});
}
events.push(ChatEvent::MessageEnded {
session_id: id.clone(),
message: normalized.clone(),
});
match normalized.stop_reason {
Some(ChatStopReason::Error) => {
cache.terminal_status = Some(ChatStatus::Error);
cache.state.status = ChatStatus::Error;
cache.state.is_streaming = false;
let message = normalized
.error_message
.clone()
.unwrap_or_else(|| "Pi provider failed without an error message".into());
events.push(status_event(&id, ChatStatus::Error));
events.push(ChatEvent::Error {
session_id: id,
code: "provider_error".into(),
message,
recoverable: true,
});
}
Some(ChatStopReason::Aborted) => {
cache.terminal_status = Some(ChatStatus::Settled);
cache.state.status = ChatStatus::Settled;
cache.state.is_streaming = false;
if !active_prompt_canceled_before_start(&cache) {
events.push(status_event(&id, ChatStatus::Settled));
}
}
_ => {}
}
Ok(events)
}
PiRpcEvent::ToolExecutionStart {
tool_call_id,
tool_name,
args,
} => {
let execution = ChatToolExecution {
id: tool_call_id.clone(),
name: tool_name,
arguments: args,
output: None,
is_error: false,
status: ChatToolStatus::Running,
};
session
.cache
.lock()
.await
.tools
.insert(tool_call_id, execution.clone());
Ok(vec![ChatEvent::ToolExecutionChanged {
session_id: id,
execution,
output_is_accumulated: false,
}])
}
PiRpcEvent::ToolExecutionUpdate {
tool_call_id,
tool_name,
args,
partial_result,
} => {
let mut cache = session.cache.lock().await;
let execution = cache.tools.get_mut(&tool_call_id).ok_or_else(|| {
invalid_event(format!("tool update has no start event: {tool_call_id}"))
})?;
if execution.name != tool_name || execution.arguments != args {
return Err(invalid_event(format!(
"tool update metadata changed for {tool_call_id}"
)));
}
execution.output = Some(partial_result);
Ok(vec![ChatEvent::ToolExecutionChanged {
session_id: id,
execution: execution.clone(),
output_is_accumulated: true,
}])
}
PiRpcEvent::ToolExecutionEnd {
tool_call_id,
tool_name,
result,
is_error,
} => {
let mut cache = session.cache.lock().await;
let mut execution = cache.tools.remove(&tool_call_id).ok_or_else(|| {
invalid_event(format!("tool end has no start event: {tool_call_id}"))
})?;
if execution.name != tool_name {
return Err(invalid_event(format!(
"tool end name changed for {tool_call_id}"
)));
}
execution.output = Some(result);
execution.is_error = is_error;
execution.status = if is_error {
ChatToolStatus::Failed
} else {
ChatToolStatus::Succeeded
};
Ok(vec![ChatEvent::ToolExecutionChanged {
session_id: id,
execution,
output_is_accumulated: false,
}])
}
PiRpcEvent::QueueUpdate {
steering,
follow_up,
} => {
session.cache.lock().await.state.pending_message_count =
(steering.len() + follow_up.len()) as u64;
Ok(vec![ChatEvent::QueueUpdated {
session_id: id,
steering,
follow_up,
}])
}
PiRpcEvent::CompactionStart { reason } => {
set_status(session, ChatStatus::Compacting, false, true).await;
Ok(vec![
status_event(&id, ChatStatus::Compacting),
ChatEvent::CompactionChanged {
session_id: id,
started: true,
reason,
result: None,
aborted: false,
will_retry: false,
error: None,
},
])
}
PiRpcEvent::CompactionEnd {
reason,
result,
aborted,
will_retry,
error_message,
} => {
let status = if will_retry {
ChatStatus::Retrying
} else {
ChatStatus::Idle
};
set_status(session, status, false, false).await;
Ok(vec![
ChatEvent::CompactionChanged {
session_id: id.clone(),
started: false,
reason,
result: result.map(|result| {
json!({
"summary": result.summary,
"firstKeptEntryId": result.first_kept_entry_id,
"tokensBefore": result.tokens_before,
"estimatedTokensAfter": result.estimated_tokens_after,
"details": result.details,
})
}),
aborted,
will_retry,
error: error_message,
},
status_event(&id, status),
])
}
PiRpcEvent::AutoRetryStart {
attempt,
max_attempts,
delay_ms,
error_message,
} => {
set_status(session, ChatStatus::Retrying, false, false).await;
Ok(vec![
status_event(&id, ChatStatus::Retrying),
ChatEvent::RetryChanged {
session_id: id,
started: true,
attempt,
max_attempts: Some(max_attempts),
delay_ms: Some(delay_ms),
success: None,
error: Some(error_message),
},
])
}
PiRpcEvent::AutoRetryEnd {
success,
attempt,
final_error,
} => {
let status = if success {
ChatStatus::Streaming
} else {
ChatStatus::Error
};
session.cache.lock().await.terminal_status = if success {
None
} else {
Some(ChatStatus::Error)
};
set_status(session, status, success, false).await;
Ok(vec![
ChatEvent::RetryChanged {
session_id: id.clone(),
started: false,
attempt,
max_attempts: None,
delay_ms: None,
success: Some(success),
error: final_error,
},
status_event(&id, status),
])
}
PiRpcEvent::ExtensionError {
extension_path,
event,
error,
} => Ok(vec![ChatEvent::ExtensionError {
session_id: id,
extension_path,
event,
error,
}]),
PiRpcEvent::ExtensionUiRequest {
id: request_id,
method,
payload,
} => {
let payload_value = Value::Object(payload.clone());
let record = ChatExtensionUiRecord {
id: request_id.clone(),
method: method.clone(),
payload: payload_value.clone(),
};
let mut cache = session.cache.lock().await;
match method.as_str() {
"notify" => push_bounded(
&mut cache.extension_notifications,
record,
MAX_EXTENSION_NOTIFICATIONS,
),
"setStatus" | "status" => {
let key = extension_key(&payload, "statusKey", &request_id);
if extension_text(&payload, "statusText")
.is_some_and(|text| !text.trim().is_empty())
{
upsert_keyed(
&mut cache.extension_statuses,
ChatExtensionKeyedRecord {
key,
id: request_id.clone(),
method: method.clone(),
payload: payload_value.clone(),
},
MAX_EXTENSION_STATUSES,
);
} else {
cache.extension_statuses.retain(|status| status.key != key);
}
}
"setWidget" => {
let key = extension_key(&payload, "widgetKey", &request_id);
let has_lines = payload
.get("widgetLines")
.and_then(Value::as_array)
.is_some_and(|lines| lines.iter().any(Value::is_string));
if has_lines {
upsert_keyed(
&mut cache.extension_widgets,
ChatExtensionKeyedRecord {
key,
id: request_id.clone(),
method: method.clone(),
payload: payload_value.clone(),
},
MAX_EXTENSION_WIDGETS,
);
} else {
cache.extension_widgets.retain(|widget| widget.key != key);
}
}
"setTitle" => {
let title = extension_text(&payload, "title")
.or_else(|| extension_text(&payload, "text"))
.map(str::to_owned);
cache.extension_title = title.clone();
cache.state.name = title;
}
"set_editor_text" => {
cache.extension_editor_text = extension_text(&payload, "text")
.or_else(|| extension_text(&payload, "value"))
.map(str::to_owned);
}
_ => {
let deadline =
payload
.get("timeout")
.and_then(Value::as_u64)
.and_then(|milliseconds| {
Instant::now().checked_add(Duration::from_millis(milliseconds))
});
upsert_pending_extension(
&mut cache.pending_extension_requests,
PendingExtensionRequest { record, deadline },
);
}
}
drop(cache);
Ok(vec![ChatEvent::ExtensionUiRequest {
session_id: id,
id: request_id,
method,
payload: payload_value,
}])
}
PiRpcEvent::ThinkingLevelChanged { level } => {
session.cache.lock().await.state.thinking_level = level;
Ok(vec![ChatEvent::ThinkingLevelChanged {
session_id: id,
level,
}])
}
PiRpcEvent::SessionInfoChanged { name } => {
let mut cache = session.cache.lock().await;
cache.extension_title = None;
cache.state.name = name.clone();
Ok(vec![ChatEvent::SessionInfoChanged {
session_id: id,
name,
}])
}
}
}
async fn normalize_message_update(
session: &PiChatSession,
message: Value,
update: Value,
) -> AgentResult<Vec<ChatEvent>> {
let object = update
.as_object()
.ok_or_else(|| invalid_event("assistantMessageEvent must be an object"))?;
let kind = required_update_string(object, "type")?;
let mut cache = session.cache.lock().await;
let message_id = stream_message_id(&mut cache, &session.id, &message, false)?;
let session_id = session.id.clone();
match kind {
"start" => {
require_update_field(object, "partial")?;
Ok(Vec::new())
}
"text_start" | "thinking_start" => {
let _ = required_content_index(object)?;
require_update_field(object, "partial")?;
Ok(Vec::new())
}
"toolcall_start" => {
let _ = required_content_index(object)?;
let partial = require_update_field(object, "partial")?;
update_live_message_from_partial(&mut cache, &message_id, &message, partial)?;
Ok(Vec::new())
}
"text_end" | "thinking_end" => {
let _ = required_content_index(object)?;
let _ = required_update_string(object, "content")?;
require_update_field(object, "partial")?;
Ok(Vec::new())
}
"text_delta" | "thinking_delta" => {
let content_index = required_content_index(object)?;
let delta = required_update_string(object, "delta")?.to_owned();
require_update_field(object, "partial")?;
let delta_kind = if kind == "text_delta" {
ChatDeltaKind::Text
} else {
ChatDeltaKind::Thinking
};
apply_message_delta(&mut cache, &message_id, content_index, delta_kind, &delta);
Ok(vec![ChatEvent::ContentDelta {
session_id,
message_id,
content_index,
kind: delta_kind,
delta,
}])
}
"toolcall_delta" => {
let partial = require_update_field(object, "partial")?;
update_live_message_from_partial(&mut cache, &message_id, &message, partial)?;
Ok(vec![ChatEvent::ToolCallDelta {
session_id,
message_id,
content_index: required_content_index(object)?,
delta: required_update_string(object, "delta")?.to_owned(),
}])
}
"toolcall_end" => {
let content_index = required_content_index(object)?;
require_update_field(object, "partial")?;
let tool_call = object
.get("toolCall")
.ok_or_else(|| invalid_event("toolcall_end requires toolCall"))?;
let tool_call = normalize_content(tool_call.clone());
if !matches!(tool_call, ChatContent::ToolCall { .. }) {
return Err(invalid_event("toolcall_end toolCall is malformed"));
}
if let Some(live) = cache.live_messages.get_mut(&message_id) {
set_content_block(&mut live.content, content_index, tool_call.clone());
}
if let Some(live) = cache.live_messages.get(&message_id).cloned() {
record_recent_live_message(&mut cache, live);
}
Ok(vec![ChatEvent::ToolCallCompleted {
session_id,
message_id,
content_index,
tool_call,
}])
}
"done" => {
let _ = required_update_string(object, "reason")?;
Ok(usage_event(&session_id, &message_id, &message)
.into_iter()
.collect())
}
"error" => {
let reason = required_update_string(object, "reason")?;
if !matches!(reason, "aborted" | "error") {
return Err(invalid_event(format!(
"assistantMessageEvent error has unsupported reason: {reason}"
)));
}
let error_message = require_update_field(object, "error")?;
if error_message.get("role").and_then(Value::as_str) != Some("assistant") {
return Err(invalid_event(
"assistantMessageEvent error must contain an assistant message",
));
}
update_live_message_from_partial(&mut cache, &message_id, &message, error_message)?;
Ok(Vec::new())
}
other => Err(invalid_event(format!(
"unsupported assistantMessageEvent type: {other}"
))),
}
}
fn update_live_message_from_partial(
cache: &mut SessionCache,
message_id: &str,
message: &Value,
partial: &Value,
) -> AgentResult<()> {
let source = if partial.get("role").is_some() {
partial.clone()
} else {
message.clone()
};
let normalized = normalize_message(message_id.to_owned(), source);
if normalized.role != ChatRole::Assistant {
return Err(invalid_event(
"toolcall partial must contain an assistant message",
));
}
cache.last_assistant_message = Some(normalized.clone());
cache
.live_messages
.insert(message_id.to_owned(), normalized.clone());
record_recent_live_message(cache, normalized);
Ok(())
}
fn required_update_string<'a>(object: &'a Map<String, Value>, field: &str) -> AgentResult<&'a str> {
object
.get(field)
.and_then(Value::as_str)
.ok_or_else(|| invalid_event(format!("assistantMessageEvent requires string {field}")))
}
fn require_update_field<'a>(object: &'a Map<String, Value>, field: &str) -> AgentResult<&'a Value> {
object
.get(field)
.ok_or_else(|| invalid_event(format!("assistantMessageEvent requires {field}")))
}
fn required_content_index(object: &Map<String, Value>) -> AgentResult<usize> {
let index = object
.get("contentIndex")
.and_then(Value::as_u64)
.ok_or_else(|| invalid_event("assistantMessageEvent requires numeric contentIndex"))?;
usize::try_from(index)
.map_err(|_| invalid_event("assistantMessageEvent contentIndex is too large"))
}
fn stream_message_id(
cache: &mut SessionCache,
session_id: &str,
message: &Value,
starting: bool,
) -> AgentResult<String> {
if starting
&& message.get("role").and_then(Value::as_str) == Some("user")
&& let Some(text) = message_text(message)
&& let Some(position) = cache
.pending_user_messages
.iter()
.position(|(pending, _)| pending == text)
&& let Some((_, id)) = cache.pending_user_messages.remove(position)
{
if let Some(pi_id) = message.get("id").and_then(Value::as_str) {
if cache.pi_message_ids.len() >= MAX_STABLE_MESSAGE_IDS {
cache.pi_message_ids.clear();
}
cache.pi_message_ids.insert(pi_id.into(), id.clone());
}
cache.active_message_id = Some(id.clone());
return Ok(id);
}
if let Some(pi_id) = message.get("id").and_then(Value::as_str) {
if let Some(id) = cache.pi_message_ids.get(pi_id) {
cache.active_message_id = Some(id.clone());
return Ok(id.clone());
}
if !starting && let Some(id) = cache.active_message_id.clone() {
if cache.pi_message_ids.len() >= MAX_STABLE_MESSAGE_IDS {
cache.pi_message_ids.clear();
}
cache.pi_message_ids.insert(pi_id.into(), id.clone());
return Ok(id);
}
let id = next_message_id(cache, session_id);
if cache.pi_message_ids.len() >= MAX_STABLE_MESSAGE_IDS {
cache.pi_message_ids.clear();
}
cache.pi_message_ids.insert(pi_id.into(), id.clone());
cache.active_message_id = Some(id.clone());
return Ok(id);
}
if starting {
let id = next_message_id(cache, session_id);
cache.active_message_id = Some(id.clone());
return Ok(id);
}
cache
.active_message_id
.clone()
.ok_or_else(|| invalid_event("message update/end has no active message"))
}
fn next_message_id(cache: &mut SessionCache, session_id: &str) -> String {
let id = format!("{session_id}:message:{}", cache.next_message_id);
cache.next_message_id = cache.next_message_id.saturating_add(1);
id
}
fn apply_message_delta(
cache: &mut SessionCache,
message_id: &str,
content_index: usize,
kind: ChatDeltaKind,
delta: &str,
) {
let Some(message) = cache.live_messages.get_mut(message_id) else {
return;
};
while message.content.len() <= content_index {
message
.content
.push(ChatContent::Unknown { value: Value::Null });
}
let slot = &mut message.content[content_index];
match (kind, slot) {
(ChatDeltaKind::Text, ChatContent::Text { text })
| (ChatDeltaKind::Thinking, ChatContent::Thinking { text }) => text.push_str(delta),
(ChatDeltaKind::Text, slot) => {
*slot = ChatContent::Text {
text: delta.to_owned(),
};
}
(ChatDeltaKind::Thinking, slot) => {
*slot = ChatContent::Thinking {
text: delta.to_owned(),
};
}
}
let updated = message.clone();
record_recent_live_message(cache, updated);
}
fn set_content_block(content: &mut Vec<ChatContent>, index: usize, block: ChatContent) {
while content.len() <= index {
content.push(ChatContent::Unknown { value: Value::Null });
}
content[index] = block;
}
fn extension_text<'a>(payload: &'a Map<String, Value>, field: &str) -> Option<&'a str> {
payload.get(field).and_then(Value::as_str)
}
fn extension_key(payload: &Map<String, Value>, field: &str, fallback: &str) -> String {
extension_text(payload, field)
.filter(|key| !key.is_empty())
.unwrap_or(fallback)
.to_owned()
}
fn push_bounded<T>(items: &mut VecDeque<T>, item: T, limit: usize) {
if items.len() == limit {
items.pop_front();
}
items.push_back(item);
}
fn slash_command_token(text: &str) -> Option<&str> {
let token = text
.starts_with('/')
.then(|| text.split_whitespace().next())??;
token
.strip_prefix('/')
.filter(|command| !command.is_empty())
}
fn sync_prompt_pending(cache: &mut SessionCache) {
cache.prompt_pending = cache.active_prompt.is_some();
}
fn clear_prompt_lifecycle(cache: &mut SessionCache, request_id: &str) -> bool {
if cache
.active_prompt
.as_ref()
.is_none_or(|active| active.request_id != request_id)
{
return false;
}
cache.active_prompt = None;
sync_prompt_pending(cache);
true
}
fn clear_active_prompt_lifecycle(cache: &mut SessionCache) -> Option<ActivePromptLifecycle> {
let active = cache.active_prompt.take();
sync_prompt_pending(cache);
active
}
fn active_prompt_canceled_before_start(cache: &SessionCache) -> bool {
cache
.active_prompt
.as_ref()
.is_some_and(|active| active.canceled_before_start)
}
fn record_cancellation_intent(cache: &mut SessionCache, request_id: &str) -> bool {
if !cache.cancellation_intents.insert(request_id.to_owned()) {
return false;
}
if cache.cancellation_intent_order.len() == MAX_PROMPT_CANCELLATION_INTENTS
&& let Some(expired) = cache.cancellation_intent_order.pop_front()
{
cache.cancellation_intents.remove(&expired);
}
cache
.cancellation_intent_order
.push_back(request_id.to_owned());
true
}
fn take_cancellation_intent(cache: &mut SessionCache, request_id: &str) -> bool {
if !cache.cancellation_intents.remove(request_id) {
return false;
}
cache
.cancellation_intent_order
.retain(|intent| intent != request_id);
true
}
fn remove_cancellation_intent(cache: &mut SessionCache, request_id: &str) {
cache.cancellation_intents.remove(request_id);
cache
.cancellation_intent_order
.retain(|intent| intent != request_id);
}
fn remember_canceled_message(cache: &mut SessionCache, message_id: &str) {
cache
.canceled_message_ids
.retain(|canceled| canceled != message_id);
if cache.canceled_message_ids.len() == MAX_STABLE_MESSAGE_IDS
&& let Some(evicted) = cache.canceled_message_ids.pop_front()
{
release_cancellation_owned_message(cache, &evicted);
}
cache.canceled_message_ids.push_back(message_id.to_owned());
}
fn mark_canceled_message(cache: &mut SessionCache, message_id: &str) {
cache
.pending_user_messages
.retain(|(_, pending_id)| pending_id != message_id);
let recent_message = cache
.recent_live_messages
.iter_mut()
.find(|message| message.id == message_id)
.map(|message| {
message.error_message = Some("Canceled".into());
message.clone()
});
if let Some(message) = cache.live_messages.get_mut(message_id) {
message.error_message = Some("Canceled".into());
} else if let Some(message) = recent_message {
cache.live_messages.insert(message_id.into(), message);
}
remember_canceled_message(cache, message_id);
}
fn release_cancellation_owned_message(cache: &mut SessionCache, message_id: &str) {
let is_pending = cache
.pending_user_messages
.iter()
.any(|(_, pending_id)| pending_id == message_id);
if is_pending || cache.active_message_id.as_deref() == Some(message_id) {
return;
}
if cache.live_messages.remove(message_id).is_none() {
return;
}
cache
.recent_live_messages
.retain(|message| message.id != message_id);
cache
.snapshot_message_ids
.retain(|_, mapped_id| mapped_id != message_id);
cache
.pi_message_ids
.retain(|_, mapped_id| mapped_id != message_id);
}
fn upsert_pending_extension(
pending: &mut VecDeque<PendingExtensionRequest>,
request: PendingExtensionRequest,
) {
if let Some(existing) = pending
.iter_mut()
.find(|existing| existing.record.id == request.record.id)
{
*existing = request;
return;
}
push_bounded(pending, request, MAX_PENDING_EXTENSION_REQUESTS);
}
fn upsert_keyed(
entries: &mut VecDeque<ChatExtensionKeyedRecord>,
record: ChatExtensionKeyedRecord,
limit: usize,
) {
if let Some(existing) = entries
.iter_mut()
.find(|existing| existing.key == record.key)
{
*existing = record;
return;
}
push_bounded(entries, record, limit);
}
async fn set_status(
session: &PiChatSession,
status: ChatStatus,
streaming: bool,
compacting: bool,
) {
let mut cache = session.cache.lock().await;
cache.state.status = status;
cache.state.is_streaming = streaming;
cache.state.is_compacting = compacting;
}
fn status_event(session_id: &str, status: ChatStatus) -> ChatEvent {
ChatEvent::StatusChanged {
session_id: session_id.into(),
status,
}
}
fn usage_event(session_id: &str, message_id: &str, message: &Value) -> Option<ChatEvent> {
message
.get("usage")
.cloned()
.map(|usage| ChatEvent::UsageUpdated {
session_id: session_id.into(),
message_id: message_id.into(),
usage,
})
}
fn normalize_snapshot_messages(
session_id: &str,
cache: &mut SessionCache,
messages: Vec<Value>,
) -> Vec<ChatMessage> {
let message_count = messages.len();
let retain_from = message_count.saturating_sub(MAX_STABLE_MESSAGE_IDS);
let last_assistant = messages
.iter()
.rposition(|message| message.get("role").and_then(Value::as_str) == Some("assistant"));
let active_snapshot_index = cache
.active_message_id
.as_ref()
.and_then(|id| cache.live_messages.get(id))
.and_then(|live| {
messages
.iter()
.enumerate()
.rev()
.find(|(_, message)| message_matches_active(message, live))
.map(|(index, _)| index)
});
let mut occurrences = HashMap::<String, usize>::new();
let mut current_snapshot_ids = HashMap::new();
let mut current_pi_ids = HashMap::new();
let mut normalized = Vec::with_capacity(
messages.len() + cache.pending_user_messages.len() + cache.canceled_message_ids.len(),
);
for (index, message) in messages.into_iter().enumerate() {
let identity = message_identity(&message);
let occurrence = occurrences.entry(identity.clone()).or_default();
let key = identity_key(&identity, *occurrence);
*occurrence += 1;
let pi_id = message.get("id").and_then(Value::as_str).map(str::to_owned);
let optimistic_user = if message.get("role").and_then(Value::as_str) == Some("user") {
message_text(&message).and_then(|text| {
cache
.pending_user_messages
.iter()
.position(|(pending, _)| pending == text)
.and_then(|position| cache.pending_user_messages.remove(position))
.map(|(_, id)| id)
})
} else {
None
};
let active_stream = (Some(index) == active_snapshot_index)
.then(|| cache.active_message_id.clone())
.flatten();
let recent_assistant = if Some(index) == last_assistant
&& message.get("role").and_then(Value::as_str) == Some("assistant")
{
cache
.last_assistant_message
.as_ref()
.filter(|recent| message_matches_cached(&message, recent))
.map(|recent| recent.id.clone())
} else {
None
};
let message_id = optimistic_user
.or_else(|| {
pi_id
.as_ref()
.and_then(|id| cache.pi_message_ids.get(id).cloned())
})
.or(active_stream)
.or(recent_assistant)
.or_else(|| cache.snapshot_message_ids.get(&key).cloned())
.unwrap_or_else(|| deterministic_snapshot_id(session_id, &key));
if index >= retain_from {
current_snapshot_ids.insert(key, message_id.clone());
if let Some(pi_id) = pi_id {
current_pi_ids.insert(pi_id, message_id.clone());
}
}
let mut authoritative = normalize_message(message_id.clone(), message);
if cache.active_message_id.as_deref() == Some(&message_id) {
if let Some(live) = cache.live_messages.get(&message_id) {
authoritative = live.clone();
}
} else {
if let Some(recent) = cache.last_assistant_message.as_ref()
&& recent.id == message_id
{
authoritative = recent.clone();
}
if authoritative.role == ChatRole::User {
cache.live_messages.remove(&message_id);
}
}
if cache.canceled_message_ids.contains(&message_id) {
authoritative.error_message = Some("Canceled".into());
}
normalized.push(authoritative);
}
let mut present = normalized
.iter()
.map(|message| message.id.clone())
.collect::<HashSet<_>>();
for (_, id) in &cache.pending_user_messages {
if !present.contains(id)
&& let Some(message) = cache.live_messages.get(id)
{
normalized.push(message.clone());
present.insert(id.clone());
}
}
for id in &cache.canceled_message_ids {
if !present.contains(id)
&& let Some(message) = cache.live_messages.get(id)
{
normalized.push(message.clone());
present.insert(id.clone());
}
}
if let Some(active) = cache.active_message_id.as_ref()
&& !normalized.iter().any(|message| &message.id == active)
&& let Some(message) = cache.live_messages.get(active)
{
normalized.push(message.clone());
}
let mut protected = cache
.pending_user_messages
.iter()
.map(|(_, id)| id.clone())
.collect::<HashSet<_>>();
protected.extend(cache.live_messages.keys().cloned());
protected.extend(cache.active_message_id.iter().cloned());
protected.extend(cache.canceled_message_ids.iter().cloned());
preserve_protected_mappings(
&cache.snapshot_message_ids,
&mut current_snapshot_ids,
&protected,
);
preserve_protected_mappings(&cache.pi_message_ids, &mut current_pi_ids, &protected);
cache.snapshot_message_ids = current_snapshot_ids;
cache.pi_message_ids = current_pi_ids;
normalized
}
fn preserve_protected_mappings(
previous: &HashMap<String, String>,
current: &mut HashMap<String, String>,
protected: &HashSet<String>,
) {
for (key, value) in previous {
if !protected.contains(value) || current.get(key) == Some(value) {
continue;
}
if !current.contains_key(key) && current.len() >= MAX_STABLE_MESSAGE_IDS {
let removable = current
.iter()
.find(|(_, mapped)| !protected.contains(*mapped))
.map(|(key, _)| key.clone());
let Some(removable) = removable else {
continue;
};
current.remove(&removable);
}
current.insert(key.clone(), value.clone());
}
}
fn message_matches_cached(message: &Value, cached: &ChatMessage) -> bool {
if cached.role != ChatRole::Assistant
|| message.get("role").and_then(Value::as_str) != Some("assistant")
{
return false;
}
message_matches_active(message, cached)
}
fn message_matches_active(message: &Value, cached: &ChatMessage) -> bool {
let role = match message.get("role").and_then(Value::as_str) {
Some("user") => ChatRole::User,
Some("assistant") => ChatRole::Assistant,
Some("toolResult") => ChatRole::Tool,
Some("bashExecution" | "system") => ChatRole::System,
_ => return false,
};
if role != cached.role {
return false;
}
if let (Some(left), Some(right)) = (
message.get("timestamp").and_then(Value::as_i64),
cached.timestamp,
) && left != right
{
return false;
}
for (field, cached_value) in [("provider", &cached.provider), ("model", &cached.model)] {
if let (Some(left), Some(right)) = (
message.get(field).and_then(Value::as_str),
cached_value.as_deref(),
) && left != right
{
return false;
}
}
let candidate = normalize_message(String::new(), message.clone());
candidate.content.len() == cached.content.len()
&& candidate
.content
.iter()
.zip(&cached.content)
.all(|(candidate, live)| content_blocks_compatible(candidate, live, &role))
}
fn content_blocks_compatible(candidate: &ChatContent, live: &ChatContent, role: &ChatRole) -> bool {
match (candidate, live) {
(ChatContent::Text { text: left }, ChatContent::Text { text: right })
| (ChatContent::Thinking { text: left }, ChatContent::Thinking { text: right })
if role == &ChatRole::Assistant =>
{
left.starts_with(right) || right.starts_with(left)
}
_ => candidate == live,
}
}
fn register_stream_message_id(cache: &mut SessionCache, identity: &str, message_id: &str) {
if cache.snapshot_message_ids.len() >= MAX_STABLE_MESSAGE_IDS {
cache.snapshot_message_ids.clear();
}
let mut occurrence = 0;
while cache
.snapshot_message_ids
.contains_key(&identity_key(identity, occurrence))
{
occurrence = occurrence.saturating_add(1);
}
cache
.snapshot_message_ids
.insert(identity_key(identity, occurrence), message_id.to_owned());
}
fn message_text(message: &Value) -> Option<&str> {
match message.get("content")? {
Value::String(text) => Some(text),
Value::Array(blocks) if blocks.len() == 1 => blocks[0].get("text")?.as_str(),
_ => None,
}
}
fn message_identity(message: &Value) -> String {
if let Some(id) = message.get("id").and_then(Value::as_str) {
return format!("pi:{id}");
}
let encoded = serde_json::to_vec(message).unwrap_or_default();
digest_hex(&encoded)
}
fn identity_key(identity: &str, occurrence: usize) -> String {
format!("{identity}:{occurrence}")
}
fn deterministic_snapshot_id(session_id: &str, identity: &str) -> String {
format!("{session_id}:snapshot:{}", digest_hex(identity.as_bytes()))
}
fn digest_hex(bytes: &[u8]) -> String {
let digest = Sha256::digest(bytes);
let mut hex = String::with_capacity(digest.len() * 2);
for byte in digest {
let _ = write!(hex, "{byte:02x}");
}
hex
}
fn normalize_message(id: String, message: Value) -> ChatMessage {
let Some(object) = message.as_object() else {
return unknown_message(id, message);
};
let role = match object.get("role").and_then(Value::as_str) {
Some("user") => ChatRole::User,
Some("assistant") => ChatRole::Assistant,
Some("toolResult") => ChatRole::Tool,
Some("bashExecution" | "system") => ChatRole::System,
_ => return unknown_message(id, message),
};
let content = match role {
ChatRole::Tool => vec![ChatContent::ToolResult {
tool_call_id: object
.get("toolCallId")
.and_then(Value::as_str)
.unwrap_or_default()
.into(),
content: object.get("content").cloned().unwrap_or(Value::Null),
is_error: object
.get("isError")
.and_then(Value::as_bool)
.unwrap_or(false),
}],
ChatRole::System if object.get("role").and_then(Value::as_str) == Some("bashExecution") => {
vec![ChatContent::Text {
text: format!(
"$ {}\n{}",
object
.get("command")
.and_then(Value::as_str)
.unwrap_or_default(),
object
.get("output")
.and_then(Value::as_str)
.unwrap_or_default()
),
}]
}
_ => normalize_message_content(object.get("content")),
};
ChatMessage {
id,
role,
content,
timestamp: object.get("timestamp").and_then(Value::as_i64),
provider: object
.get("provider")
.and_then(Value::as_str)
.map(str::to_owned),
model: object
.get("model")
.and_then(Value::as_str)
.map(str::to_owned),
usage: object.get("usage").cloned(),
stop_reason: object
.get("stopReason")
.and_then(Value::as_str)
.map(parse_stop_reason),
error_message: object
.get("errorMessage")
.and_then(Value::as_str)
.map(str::to_owned),
}
}
fn unknown_message(id: String, value: Value) -> ChatMessage {
ChatMessage {
id,
role: ChatRole::System,
content: vec![ChatContent::Unknown { value }],
timestamp: None,
provider: None,
model: None,
usage: None,
stop_reason: None,
error_message: None,
}
}
fn parse_stop_reason(reason: &str) -> ChatStopReason {
match reason {
"stop" => ChatStopReason::Stop,
"length" => ChatStopReason::Length,
"toolUse" => ChatStopReason::ToolUse,
"error" => ChatStopReason::Error,
"aborted" => ChatStopReason::Aborted,
other => ChatStopReason::Unknown(other.to_owned()),
}
}
fn normalize_message_content(content: Option<&Value>) -> Vec<ChatContent> {
match content {
Some(Value::String(text)) => vec![ChatContent::Text { text: text.clone() }],
Some(Value::Array(blocks)) => blocks.iter().cloned().map(normalize_content).collect(),
Some(value) => vec![ChatContent::Unknown {
value: value.clone(),
}],
None => Vec::new(),
}
}
fn normalize_content(value: Value) -> ChatContent {
let Some(object) = value.as_object() else {
return ChatContent::Unknown { value };
};
match object.get("type").and_then(Value::as_str) {
Some("text") => object
.get("text")
.and_then(Value::as_str)
.map(|text| ChatContent::Text { text: text.into() })
.unwrap_or(ChatContent::Unknown { value }),
Some("thinking") => object
.get("thinking")
.and_then(Value::as_str)
.map(|text| ChatContent::Thinking { text: text.into() })
.unwrap_or(ChatContent::Unknown { value }),
Some("toolCall") => match (
object.get("id").and_then(Value::as_str),
object.get("name").and_then(Value::as_str),
object.get("arguments"),
) {
(Some(id), Some(name), Some(arguments)) => ChatContent::ToolCall {
id: id.into(),
name: name.into(),
arguments: arguments.clone(),
},
_ => ChatContent::Unknown { value },
},
_ => ChatContent::Unknown { value },
}
}
fn model_from_value(value: Value, fallback: Option<(&str, &str)>) -> Option<ChatModel> {
let provider = value
.get("provider")
.and_then(Value::as_str)
.map(str::to_owned)
.or_else(|| fallback.map(|(provider, _)| provider.to_owned()))?;
let id = value
.get("id")
.and_then(Value::as_str)
.map(str::to_owned)
.or_else(|| fallback.map(|(_, id)| id.to_owned()))?;
Some(ChatModel {
provider,
id,
value,
})
}
fn status_from_state(state: &PiSessionState) -> ChatStatus {
if state.is_compacting {
ChatStatus::Compacting
} else if state.is_streaming {
ChatStatus::Streaming
} else {
ChatStatus::Idle
}
}
fn expect_state(response: PiRpcResponse) -> AgentResult<PiSessionState> {
match response {
PiRpcResponse::GetState { data, .. } => Ok(data),
response => Err(unexpected_response(response, "get_state")),
}
}
fn expect_messages(response: PiRpcResponse) -> AgentResult<PiMessagesData> {
match response {
PiRpcResponse::GetMessages { data, .. } => Ok(data),
response => Err(unexpected_response(response, "get_messages")),
}
}
fn expect_stats(response: PiRpcResponse) -> AgentResult<PiSessionStats> {
match response {
PiRpcResponse::GetSessionStats { data, .. } => Ok(data),
response => Err(unexpected_response(response, "get_session_stats")),
}
}
fn expect_set_model(response: PiRpcResponse) -> AgentResult<Value> {
match response {
PiRpcResponse::SetModel { data, .. } => Ok(data),
response => Err(unexpected_response(response, "set_model")),
}
}
fn expect_command(response: PiRpcResponse, expected: &str) -> AgentResult<()> {
if response.command() == expected && response.success() {
Ok(())
} else {
Err(unexpected_response(response, expected))
}
}
fn unexpected_response(response: PiRpcResponse, expected: &str) -> AgentError {
if let Some(error) = response.error() {
AgentError::new(
ErrorCode::InvalidMessage,
format!("Pi RPC {expected} failed: {error}"),
)
} else {
AgentError::new(
ErrorCode::InvalidMessage,
format!(
"Pi RPC returned {} while waiting for {expected}",
response.command()
),
)
}
}
fn command_discovery_failed(error: AgentError) -> AgentError {
if error.code() == ErrorCode::SessionNotFound {
error
} else {
AgentError::new(
ErrorCode::ChatCommandDiscoveryFailed,
"Pi command discovery failed",
)
}
}
fn validate_model_selection(provider: Option<&str>, model_id: Option<&str>) -> AgentResult<()> {
if provider.is_some() == model_id.is_some() {
Ok(())
} else {
Err(AgentError::new(
ErrorCode::InvalidMessage,
"provider and model id must be supplied together",
))
}
}
fn invalid_event(message: impl Into<String>) -> AgentError {
AgentError::new(ErrorCode::InvalidMessage, message)
}
fn session_missing() -> AgentError {
AgentError::new(ErrorCode::SessionNotFound, "Pi chat session not found")
}
fn session_exists() -> AgentError {
AgentError::new(
ErrorCode::SessionAlreadyExists,
"Pi chat session already exists",
)
}
fn initialization_cancelled() -> AgentError {
AgentError::new(
ErrorCode::BackendDisconnected,
"Pi chat initialization was cancelled",
)
}