use crate::assembler::{Assembled, AssemblyError, ResponseAssembler};
use crate::budget::{evaluate, BudgetPolicy, BudgetVerdict};
use crate::chain::project_and_validate;
use crate::compaction::{
apply_summary, generate_summary, select_compaction_range, source_hash, CompactionConfig,
CompactionModelSelector, SummaryPayload,
};
use crate::error::{ProviderError, RunFailureReason, ToolError};
use crate::event::{EventEnvelope, RealtimeEvent};
use crate::ids::{BranchId, ModelId, RunId, SessionId, ToolCallId, TurnId};
use crate::message::{ContentBlock, FinishReason, Message, ToolResultPayload, Usage};
use crate::policy::{
ApprovalDecision, ApprovalHandler, ApprovalRequest, Decision, Policy, PolicyRequest,
SteerSource,
};
use crate::prompt::{PromptComposer, PromptContext};
use crate::provider::{
validate_request_capabilities, CredentialProvider, GenerationOptions, ModelCapabilities,
ModelMessage, ModelProvider, ModelRequest, ProviderEvent, ProviderStream, ReasoningLevel,
TokenCounter, TokenMeasurement,
};
use crate::skill::SkillRegistry;
use crate::tool::{
Concurrency, PreparedToolCall, Tool, ToolContext, ToolExecutionContext, ToolProgress, ToolSpec,
};
use async_trait::async_trait;
use futures::{future::join_all, StreamExt};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
#[async_trait]
pub trait RunHooks: Send + Sync {
async fn on_turn_started(&self, _turn: u32) -> Result<(), RunFailureReason> {
Ok(())
}
async fn on_assistant_message(&self, _message: &Message) -> Result<(), RunFailureReason> {
Ok(())
}
async fn on_tool_calls_planned(
&self,
_calls: &[(ToolCallId, String, serde_json::Value)],
) -> Result<(), RunFailureReason> {
Ok(())
}
async fn on_tool_results(
&self,
_results: &[ToolResultPayload],
) -> Result<(), RunFailureReason> {
Ok(())
}
async fn on_turn_completed(
&self,
_turn: u32,
_finish_reason: FinishReason,
) -> Result<(), RunFailureReason> {
Ok(())
}
async fn on_steer_applied(&self, _contents: &[String]) -> Result<(), RunFailureReason> {
Ok(())
}
async fn on_summary_checkpoint(
&self,
_summary: &SummaryPayload,
) -> Result<(), RunFailureReason> {
Ok(())
}
}
pub struct NoopHooks;
impl RunHooks for NoopHooks {}
pub struct LoopParams {
pub session_id: SessionId,
pub branch_id: BranchId,
pub run_id: RunId,
pub provider: Arc<dyn ModelProvider>,
pub token_counter: Arc<dyn TokenCounter>,
pub credentials: Arc<dyn CredentialProvider>,
pub tools: Vec<Arc<dyn Tool>>,
pub model: ModelId,
pub reasoning: ReasoningLevel,
pub generation: GenerationOptions,
pub capabilities: ModelCapabilities,
pub system_prompt: String,
pub prompt: Option<Arc<PromptComposer>>,
pub skills: Option<Arc<SkillRegistry>>,
pub definition_id: String,
pub history: Vec<Message>,
pub compaction: Option<CompactionConfig>,
pub compaction_selector: Arc<dyn CompactionModelSelector>,
pub provider_options: serde_json::Value,
pub budget: BudgetPolicy,
pub events: mpsc::Sender<EventEnvelope<RealtimeEvent>>,
pub cancel: CancellationToken,
pub max_turns: Option<u32>,
pub hooks: Arc<dyn RunHooks>,
pub cancel_grace: std::time::Duration,
pub policy: Arc<dyn Policy>,
pub approval: Option<Arc<dyn ApprovalHandler>>,
pub steer: Option<Arc<dyn SteerSource>>,
}
#[derive(Debug)]
pub enum LoopOutcome {
Completed {
messages: Vec<Message>,
},
Failed {
reason: RunFailureReason,
messages: Vec<Message>,
},
Cancelled {
messages: Vec<Message>,
},
}
const CANCELLED_BEFORE_START_TEXT: &str = "工具在执行前被取消";
pub async fn run_agent_loop(params: LoopParams) -> LoopOutcome {
let seq = Arc::new(AtomicU64::new(0));
let mut produced: Vec<Message> = Vec::new();
emit(¶ms, &seq, RealtimeEvent::RunStarted).await;
if let Err(error) = project_and_validate(¶ms.history) {
tracing::warn!(?error, "chain validation failed");
return finish(
¶ms,
&seq,
LoopOutcome::Failed {
reason: RunFailureReason::Internal,
messages: produced,
},
)
.await;
}
let specs: Vec<ToolSpec> = params.tools.iter().map(|t| t.spec()).collect();
let outcome = drive_turns(¶ms, &seq, &specs, params.history.clone(), &mut produced).await;
finish(¶ms, &seq, outcome).await
}
async fn drive_turns(
params: &LoopParams,
seq: &Arc<AtomicU64>,
specs: &[ToolSpec],
mut working: Vec<Message>,
produced: &mut Vec<Message>,
) -> LoopOutcome {
let mut turn: u32 = 0;
let mut overflow_retried = false;
loop {
turn += 1;
if let Some(max) = params.max_turns {
if turn > max {
return LoopOutcome::Failed {
reason: RunFailureReason::Internal,
messages: std::mem::take(produced),
};
}
}
let mut chain = match project_and_validate(&working) {
Ok(chain) => chain,
Err(error) => {
tracing::warn!(?error, "chain validation failed");
return LoopOutcome::Failed {
reason: RunFailureReason::Internal,
messages: std::mem::take(produced),
};
}
};
let system_prompt = match ¶ms.prompt {
Some(composer) => {
let context = PromptContext {
session_id: params.session_id.clone(),
run_id: params.run_id.clone(),
turn,
definition_id: params.definition_id.clone(),
skill_catalog: params
.skills
.as_ref()
.map(|registry| registry.summaries())
.unwrap_or_default(),
};
match composer
.render_system_prompt(&context, Some(¶ms.provider.id()))
.await
{
Ok(prompt) => prompt,
Err(error) => {
return LoopOutcome::Failed {
reason: RunFailureReason::Prompt(error),
messages: std::mem::take(produced),
}
}
}
}
None => params.system_prompt.clone(),
};
let request =
match budget_phase(params, specs, &system_prompt, &mut working, &mut chain).await {
BudgetPhase::Ready { request } => request,
BudgetPhase::Cancelled => {
return LoopOutcome::Cancelled {
messages: std::mem::take(produced),
}
}
BudgetPhase::Failed(reason) => {
return LoopOutcome::Failed {
reason,
messages: std::mem::take(produced),
}
}
};
if let Err(reason) = params.hooks.on_turn_started(turn).await {
return LoopOutcome::Failed {
reason,
messages: std::mem::take(produced),
};
}
emit(params, seq, RealtimeEvent::TurnStarted { turn }).await;
let mut request = request;
let assembled = loop {
let stream = match params
.provider
.stream(
request.clone(),
params.credentials.as_ref(),
params.cancel.clone(),
)
.await
{
Ok(stream) => stream,
Err(ProviderError::Cancelled) => {
return LoopOutcome::Cancelled {
messages: std::mem::take(produced),
}
}
Err(ProviderError::ContextOverflow) => {
match overflow_retry(
params,
specs,
&system_prompt,
&mut working,
&mut chain,
&mut overflow_retried,
)
.await
{
OverflowRetry::Retried { next_request } => {
request = next_request;
continue;
}
OverflowRetry::Cancelled => {
return LoopOutcome::Cancelled {
messages: std::mem::take(produced),
}
}
OverflowRetry::Exhausted => {
return LoopOutcome::Failed {
reason: RunFailureReason::Provider(ProviderError::ContextOverflow),
messages: std::mem::take(produced),
}
}
}
}
Err(error) => {
return LoopOutcome::Failed {
reason: RunFailureReason::Provider(error),
messages: std::mem::take(produced),
}
}
};
match consume_stream(params, seq, stream).await {
StreamOutcome::Assembled(assembled) => break assembled,
StreamOutcome::Cancelled => {
return LoopOutcome::Cancelled {
messages: std::mem::take(produced),
}
}
StreamOutcome::Overflow => {
match overflow_retry(
params,
specs,
&system_prompt,
&mut working,
&mut chain,
&mut overflow_retried,
)
.await
{
OverflowRetry::Retried { next_request } => {
request = next_request;
continue;
}
OverflowRetry::Cancelled => {
return LoopOutcome::Cancelled {
messages: std::mem::take(produced),
}
}
OverflowRetry::Exhausted => {
return LoopOutcome::Failed {
reason: RunFailureReason::Provider(ProviderError::ContextOverflow),
messages: std::mem::take(produced),
}
}
}
}
StreamOutcome::Failed(error) => {
return LoopOutcome::Failed {
reason: RunFailureReason::Provider(error),
messages: std::mem::take(produced),
}
}
}
};
let (assistant, finish_reason, calls, assistant_blocks) = match assembled {
Assembled::Complete {
blocks,
finish_reason,
..
} => {
let calls = blocks
.iter()
.filter_map(|b| match b {
ContentBlock::ToolCall {
id,
name,
arguments,
} => Some((id.clone(), name.clone(), arguments.clone())),
_ => None,
})
.collect::<Vec<_>>();
(
Message::Assistant {
blocks: blocks.clone(),
finish_reason,
truncated: false,
},
finish_reason,
calls,
blocks,
)
}
Assembled::Truncated { blocks, .. } => (
Message::Assistant {
blocks: blocks.clone(),
finish_reason: FinishReason::Length,
truncated: true,
},
FinishReason::Length,
Vec::new(),
blocks,
),
Assembled::Invalid { error } => {
return LoopOutcome::Failed {
reason: RunFailureReason::Provider(ProviderError::Protocol(format!(
"{error:?}"
))),
messages: std::mem::take(produced),
}
}
};
if let Err(reason) = params.hooks.on_assistant_message(&assistant).await {
return LoopOutcome::Failed {
reason,
messages: std::mem::take(produced),
};
}
produced.push(assistant);
if let Err(reason) = params.hooks.on_turn_completed(turn, finish_reason).await {
return LoopOutcome::Failed {
reason,
messages: std::mem::take(produced),
};
}
emit(params, seq, RealtimeEvent::TurnCompleted { finish_reason }).await;
if calls.is_empty() {
working.push(Message::Assistant {
blocks: assistant_blocks,
finish_reason,
truncated: false,
});
match drain_steer(params, seq, &mut working).await {
Err(reason) => {
return LoopOutcome::Failed {
reason,
messages: std::mem::take(produced),
}
}
Ok(items) if !items.is_empty() => continue,
Ok(_) => {}
}
return LoopOutcome::Completed {
messages: std::mem::take(produced),
};
}
if let Err(reason) = params.hooks.on_tool_calls_planned(&calls).await {
return LoopOutcome::Failed {
reason,
messages: std::mem::take(produced),
};
}
let turn_id = TurnId::generate();
enum Gated {
Finished(ToolResultPayload),
Allowed(Arc<dyn Tool>, PreparedToolCall),
}
let mut gated: Vec<Gated> = Vec::with_capacity(calls.len());
for (call_id, name, arguments) in calls {
if params.cancel.is_cancelled() {
gated.push(Gated::Finished(ToolResultPayload {
call_id,
is_error: true,
text: CANCELLED_BEFORE_START_TEXT.into(),
}));
continue;
}
let Some(tool) = params.tools.iter().find(|t| t.spec().name == name) else {
gated.push(Gated::Finished(ToolResultPayload {
call_id,
is_error: true,
text: format!("unknown tool: {name}"),
}));
continue;
};
let context = ToolContext {
session_id: params.session_id.clone(),
run_id: params.run_id.clone(),
turn_id: turn_id.clone(),
call_id: call_id.clone(),
};
let prepared = match tool.prepare(arguments, &context).await {
Ok(prepared) => prepared,
Err(error) => {
gated.push(Gated::Finished(ToolResultPayload {
call_id,
is_error: true,
text: format!("prepare failed: {error}"),
}));
continue;
}
};
let decision = tokio::select! {
biased;
_ = params.cancel.cancelled() => {
gated.push(Gated::Finished(ToolResultPayload {
call_id,
is_error: true,
text: CANCELLED_BEFORE_START_TEXT.into(),
}));
continue;
}
d = params.policy.evaluate(PolicyRequest {
call: prepared.clone(),
session_id: params.session_id.clone(),
run_id: params.run_id.clone(),
turn,
}) => match d {
Ok(d) => d,
Err(e) => Decision::Deny {
reason: format!("policy error: {e}"),
},
},
};
match decision {
Decision::Allow => gated.push(Gated::Allowed(tool.clone(), prepared)),
Decision::Deny { reason } => gated.push(Gated::Finished(ToolResultPayload {
call_id,
is_error: true,
text: format!("权限拒绝: {reason}"),
})),
Decision::Ask => {
let Some(handler) = params.approval.as_ref() else {
gated.push(Gated::Finished(ToolResultPayload {
call_id,
is_error: true,
text: "权限拒绝: 审批未启用".into(),
}));
continue;
};
let approval_decision = tokio::select! {
biased;
_ = params.cancel.cancelled() => ApprovalDecision::Denied {
reason: CANCELLED_BEFORE_START_TEXT.into(),
},
d = handler.wait(ApprovalRequest {
call_id: call_id.clone(),
name: name.clone(),
capabilities: prepared.capabilities.clone(),
}, params.cancel.clone()) => d,
};
match approval_decision {
ApprovalDecision::Approved => {
gated.push(Gated::Allowed(tool.clone(), prepared))
}
ApprovalDecision::Denied { reason } => {
gated.push(Gated::Finished(ToolResultPayload {
call_id,
is_error: true,
text: format!("权限拒绝: {reason}"),
}))
}
}
}
}
}
let mut slots: Vec<Option<ToolResultPayload>> = Vec::with_capacity(gated.len());
let mut allowed: Vec<(Arc<dyn Tool>, PreparedToolCall)> = Vec::new();
let mut allowed_indices: Vec<usize> = Vec::new();
for (i, g) in gated.into_iter().enumerate() {
match g {
Gated::Finished(payload) => slots.push(Some(payload)),
Gated::Allowed(tool, prepared) => {
allowed_indices.push(i);
allowed.push((tool, prepared));
slots.push(None);
}
}
}
if allowed.len() > 1
&& allowed
.iter()
.all(|(tool, _)| tool.spec().concurrency == Concurrency::ParallelSafe)
{
let futures = allowed.iter().map(|(tool, prepared)| {
execute_prepared(params, seq, tool.clone(), prepared.clone())
});
let executed = join_all(futures).await;
for (slot, payload) in allowed_indices.into_iter().zip(executed) {
slots[slot] = Some(payload);
}
} else {
for (i, (tool, prepared)) in allowed.into_iter().enumerate() {
if params.cancel.is_cancelled() {
slots[allowed_indices[i]] = Some(ToolResultPayload {
call_id: prepared.call_id,
is_error: true,
text: CANCELLED_BEFORE_START_TEXT.into(),
});
continue;
}
slots[allowed_indices[i]] =
Some(execute_prepared(params, seq, tool, prepared).await);
}
}
let results: Vec<ToolResultPayload> = slots
.into_iter()
.map(|slot| slot.expect("every slot filled by gate or execution phase"))
.collect();
if let Err(reason) = params.hooks.on_tool_results(&results).await {
return LoopOutcome::Failed {
reason,
messages: std::mem::take(produced),
};
}
working.push(Message::Assistant {
blocks: assistant_blocks,
finish_reason,
truncated: false,
});
working.push(Message::ToolResult {
results: results.clone(),
});
produced.push(Message::ToolResult { results });
if params.cancel.is_cancelled() {
return LoopOutcome::Cancelled {
messages: std::mem::take(produced),
};
}
match drain_steer(params, seq, &mut working).await {
Err(reason) => {
return LoopOutcome::Failed {
reason,
messages: std::mem::take(produced),
}
}
Ok(items) if !items.is_empty() => continue,
Ok(_) => {}
}
}
}
enum BudgetPhase {
Ready { request: ModelRequest },
Cancelled,
Failed(RunFailureReason),
}
fn build_model_request(
params: &LoopParams,
specs: &[ToolSpec],
system_prompt: &str,
chain: &[ModelMessage],
) -> ModelRequest {
ModelRequest {
model: params.model.clone(),
system_prompt: system_prompt.to_string(),
messages: chain.to_vec(),
tools: specs.to_vec(),
reasoning: params.reasoning,
generation: params.generation.clone(),
provider_options: params.provider_options.clone(),
}
}
async fn count_input(
params: &LoopParams,
request: &ModelRequest,
) -> Result<TokenMeasurement, BudgetPhase> {
params
.token_counter
.count_input(request, params.credentials.as_ref(), params.cancel.clone())
.await
.map_err(|error| match error {
ProviderError::Cancelled => BudgetPhase::Cancelled,
other => BudgetPhase::Failed(RunFailureReason::Provider(other)),
})
}
fn exceeded_reason(params: &LoopParams, measurement: &TokenMeasurement) -> RunFailureReason {
match evaluate(¶ms.capabilities, measurement, ¶ms.budget) {
BudgetVerdict::Exceeded {
measured_tokens,
available_tokens,
source,
} => RunFailureReason::ContextBudgetExceeded {
measured_tokens,
available_tokens,
source,
},
BudgetVerdict::Within { .. } => RunFailureReason::ContextBudgetExceeded {
measured_tokens: measurement.input_tokens,
available_tokens: 0,
source: measurement.source,
},
}
}
#[allow(clippy::too_many_arguments)]
async fn budget_phase(
params: &LoopParams,
specs: &[ToolSpec],
system_prompt: &str,
working: &mut Vec<Message>,
chain: &mut Vec<ModelMessage>,
) -> BudgetPhase {
let request = build_model_request(params, specs, system_prompt, chain);
if let Err(reason) =
validate_request_capabilities(¶ms.provider.id(), &request, ¶ms.capabilities)
{
return BudgetPhase::Failed(reason);
}
let measurement = match count_input(params, &request).await {
Ok(m) => m,
Err(phase) => return phase,
};
let verdict = evaluate(¶ms.capabilities, &measurement, ¶ms.budget);
let Some(config) = params.compaction else {
return match verdict {
BudgetVerdict::Exceeded { .. } => {
BudgetPhase::Failed(exceeded_reason(params, &measurement))
}
BudgetVerdict::Within { .. } => BudgetPhase::Ready { request },
};
};
let triggered = match verdict {
BudgetVerdict::Exceeded { .. } => true,
BudgetVerdict::Within { available_tokens } => {
(measurement.input_tokens as f64) > (available_tokens as f64) * config.trigger_ratio
}
};
let originally_within = matches!(verdict, BudgetVerdict::Within { .. });
if !triggered {
return BudgetPhase::Ready { request };
}
let mut current = measurement;
let mut last_request = request;
let mut ready: Option<(ModelRequest, TokenMeasurement)> = None;
for _attempt in 0..config.max_attempts {
let Some((cover_len, retain_from)) =
select_compaction_range(working, config.keep_recent_turns)
else {
return if originally_within {
BudgetPhase::Ready {
request: last_request,
}
} else {
BudgetPhase::Failed(exceeded_reason(params, ¤t))
};
};
let summary_model = params.compaction_selector.select_model(¶ms.model);
let covered = working[..cover_len].to_vec();
let text = tokio::select! {
biased;
_ = params.cancel.cancelled() => return BudgetPhase::Cancelled,
text = generate_summary(
params.provider.as_ref(),
params.credentials.as_ref(),
params.cancel.clone(),
&summary_model,
&covered,
) => match text {
Ok(text) => text,
Err(ProviderError::Cancelled) => return BudgetPhase::Cancelled,
Err(_) => continue,
},
};
let payload = SummaryPayload {
text: text.clone(),
covered_message_count: cover_len,
retain_from,
source_hash: source_hash(&covered),
usage: Usage::default(),
};
if let Err(reason) = params.hooks.on_summary_checkpoint(&payload).await {
return BudgetPhase::Failed(reason);
}
*working = apply_summary(working, &text, retain_from);
*chain = match project_and_validate(working) {
Ok(chain) => chain,
Err(error) => {
tracing::warn!(?error, "post-compaction chain validation failed");
return BudgetPhase::Failed(RunFailureReason::Internal);
}
};
let next_request = build_model_request(params, specs, system_prompt, chain);
let next = match count_input(params, &next_request).await {
Ok(m) => m,
Err(phase) => return phase,
};
last_request = next_request.clone();
match evaluate(¶ms.capabilities, &next, ¶ms.budget) {
BudgetVerdict::Within { .. } => {
ready = Some((next_request, next));
break;
}
BudgetVerdict::Exceeded { .. } => {
if next.input_tokens < current.input_tokens {
current = next;
continue;
}
return if originally_within {
BudgetPhase::Ready {
request: last_request,
}
} else {
BudgetPhase::Failed(exceeded_reason(params, &next))
};
}
}
}
if let Some((request, _)) = ready {
return BudgetPhase::Ready { request };
}
if originally_within {
return BudgetPhase::Ready {
request: last_request,
};
}
BudgetPhase::Failed(exceeded_reason(params, ¤t))
}
enum ForcedCompaction {
Succeeded { next_request: ModelRequest },
Cancelled,
Failed,
}
enum OverflowRetry {
Retried { next_request: ModelRequest },
Cancelled,
Exhausted,
}
#[allow(clippy::too_many_arguments)]
async fn overflow_retry(
params: &LoopParams,
specs: &[ToolSpec],
system_prompt: &str,
working: &mut Vec<Message>,
chain: &mut Vec<ModelMessage>,
overflow_retried: &mut bool,
) -> OverflowRetry {
if *overflow_retried || params.compaction.is_none() {
return OverflowRetry::Exhausted;
}
*overflow_retried = true;
match forced_compaction(params, specs, system_prompt, working, chain).await {
ForcedCompaction::Succeeded { next_request } => OverflowRetry::Retried { next_request },
ForcedCompaction::Cancelled => OverflowRetry::Cancelled,
ForcedCompaction::Failed => OverflowRetry::Exhausted,
}
}
async fn forced_compaction(
params: &LoopParams,
specs: &[ToolSpec],
system_prompt: &str,
working: &mut Vec<Message>,
chain: &mut Vec<ModelMessage>,
) -> ForcedCompaction {
let Some(config) = params.compaction else {
return ForcedCompaction::Failed;
};
let Some((cover_len, retain_from)) = select_compaction_range(working, config.keep_recent_turns)
else {
return ForcedCompaction::Failed;
};
let summary_model = params.compaction_selector.select_model(¶ms.model);
let covered = working[..cover_len].to_vec();
let text = tokio::select! {
biased;
_ = params.cancel.cancelled() => return ForcedCompaction::Cancelled,
text = generate_summary(
params.provider.as_ref(),
params.credentials.as_ref(),
params.cancel.clone(),
&summary_model,
&covered,
) => match text {
Ok(text) => text,
Err(ProviderError::Cancelled) => return ForcedCompaction::Cancelled,
Err(_) => return ForcedCompaction::Failed,
},
};
let payload = SummaryPayload {
text: text.clone(),
covered_message_count: cover_len,
retain_from,
source_hash: source_hash(&covered),
usage: Usage::default(),
};
if params.hooks.on_summary_checkpoint(&payload).await.is_err() {
return ForcedCompaction::Failed;
}
*working = apply_summary(working, &text, retain_from);
*chain = match project_and_validate(working) {
Ok(chain) => chain,
Err(error) => {
tracing::warn!(?error, "post-compaction chain validation failed");
return ForcedCompaction::Failed;
}
};
let next_request = build_model_request(params, specs, system_prompt, chain);
match count_input(params, &next_request).await {
Ok(measurement) => {
if matches!(
evaluate(¶ms.capabilities, &measurement, ¶ms.budget),
BudgetVerdict::Exceeded { .. }
) {
ForcedCompaction::Failed
} else {
ForcedCompaction::Succeeded { next_request }
}
}
Err(BudgetPhase::Cancelled) => ForcedCompaction::Cancelled,
Err(_) => ForcedCompaction::Failed,
}
}
async fn drain_steer(
params: &LoopParams,
seq: &Arc<AtomicU64>,
working: &mut Vec<Message>,
) -> Result<Vec<String>, RunFailureReason> {
let Some(steer) = params.steer.as_ref() else {
return Ok(Vec::new());
};
let items = steer.pending().await;
if items.is_empty() {
return Ok(Vec::new());
}
params.hooks.on_steer_applied(&items).await?;
for item in &items {
working.push(Message::User {
blocks: vec![ContentBlock::Text { text: item.clone() }],
});
}
emit(
params,
seq,
RealtimeEvent::SteerInjected {
count: items.len() as u32,
},
)
.await;
Ok(items)
}
enum StreamOutcome {
Assembled(Assembled),
Cancelled,
Overflow,
Failed(ProviderError),
}
async fn consume_stream(
params: &LoopParams,
seq: &Arc<AtomicU64>,
mut stream: ProviderStream,
) -> StreamOutcome {
let mut assembler = ResponseAssembler::new();
let mut saw_output = false;
loop {
let item = tokio::select! {
biased;
_ = params.cancel.cancelled() => return StreamOutcome::Cancelled,
item = stream.next() => item,
};
match item {
None => {
return match assembler.finalize() {
Assembled::Invalid {
error: AssemblyError::EndedWithoutCompletion,
} => StreamOutcome::Failed(ProviderError::Protocol(
"stream ended without ResponseCompleted".into(),
)),
assembled => StreamOutcome::Assembled(assembled),
};
}
Some(Ok(event)) => {
saw_output = true;
if let Some(realtime) = to_realtime(&event) {
emit(params, seq, realtime).await;
}
if let Err(error) = assembler.push(event) {
return StreamOutcome::Failed(ProviderError::Protocol(format!("{error:?}")));
}
}
Some(Err(error)) => match error {
ProviderError::Cancelled => return StreamOutcome::Cancelled,
ProviderError::ContextOverflow if !saw_output => return StreamOutcome::Overflow,
other => return StreamOutcome::Failed(other),
},
}
}
}
fn to_realtime(event: &ProviderEvent) -> Option<RealtimeEvent> {
match event {
ProviderEvent::TextDelta { block, text } => Some(RealtimeEvent::TextDelta {
block: *block,
text: text.clone(),
}),
ProviderEvent::ReasoningDelta { block, text } => Some(RealtimeEvent::ReasoningDelta {
block: *block,
text: text.clone(),
}),
ProviderEvent::ToolCallStarted { block, id, name } => {
Some(RealtimeEvent::ToolCallStarted {
block: *block,
id: ToolCallId::from(id.clone()),
name: name.clone(),
})
}
ProviderEvent::UsageUpdated(usage) => Some(RealtimeEvent::UsageUpdated { usage: *usage }),
_ => None,
}
}
async fn execute_prepared(
params: &LoopParams,
seq: &Arc<AtomicU64>,
tool: Arc<dyn Tool>,
prepared: crate::tool::PreparedToolCall,
) -> ToolResultPayload {
let call_id = prepared.call_id.clone();
let (progress_tx, mut progress_rx) = mpsc::channel::<ToolProgress>(16);
let exec = ToolExecutionContext {
cancel: params.cancel.clone(),
progress: progress_tx,
};
let tool = tool.clone();
let join = tokio::spawn(async move { tool.execute(prepared, exec).await });
let session_id = params.session_id.clone();
let branch_id = params.branch_id.clone();
let run_id = params.run_id.clone();
let events = params.events.clone();
let drainer_seq = seq.clone();
let drainer_call_id = call_id.clone();
let drainer = tokio::spawn(async move {
while let Some(progress) = progress_rx.recv().await {
let run_seq = drainer_seq.fetch_add(1, Ordering::SeqCst) + 1;
let _ = events
.send(EventEnvelope {
session_id: session_id.clone(),
branch_id: branch_id.clone(),
run_id: Some(run_id.clone()),
revision: 0,
run_seq: Some(run_seq),
payload: RealtimeEvent::ToolProgress {
call_id: drainer_call_id.clone(),
message: progress.message,
},
})
.await;
}
});
let mut join = std::pin::pin!(join);
let result = tokio::select! {
biased;
_ = params.cancel.cancelled() => {
match tokio::time::timeout(params.cancel_grace, &mut join).await {
Ok(result) => result,
Err(_elapsed) => Ok(Err(ToolError::Execution(
"工具可能已经执行,但未记录到结果,不得假定可以安全重试".into(),
))),
}
}
result = &mut join => result,
};
let _ = tokio::time::timeout(std::time::Duration::from_secs(1), drainer).await;
match result {
Ok(Ok(output)) => ToolResultPayload {
call_id,
is_error: output.is_error,
text: output.text,
},
Ok(Err(error)) => ToolResultPayload {
call_id,
is_error: true,
text: format!("tool error: {error}"),
},
Err(join_error) => ToolResultPayload {
call_id,
is_error: true,
text: format!("tool task failed: {join_error}"),
},
}
}
async fn emit(params: &LoopParams, seq: &Arc<AtomicU64>, payload: RealtimeEvent) {
let run_seq = seq.fetch_add(1, Ordering::SeqCst) + 1;
let _ = params
.events
.send(EventEnvelope {
session_id: params.session_id.clone(),
branch_id: params.branch_id.clone(),
run_id: Some(params.run_id.clone()),
revision: 0,
run_seq: Some(run_seq),
payload,
})
.await;
}
async fn finish(params: &LoopParams, seq: &Arc<AtomicU64>, outcome: LoopOutcome) -> LoopOutcome {
let terminal = match &outcome {
LoopOutcome::Completed { .. } => RealtimeEvent::RunCompleted,
LoopOutcome::Failed { reason, .. } => RealtimeEvent::RunFailed {
reason: reason.clone(),
},
LoopOutcome::Cancelled { .. } => RealtimeEvent::RunCancelled,
};
emit(params, seq, terminal).await;
outcome
}