Harden direct provider lifecycle handling
This commit is contained in:
@@ -88,7 +88,9 @@ impl ProviderRunResponseProjector {
|
||||
})
|
||||
}
|
||||
ProviderRunProjection::ModelEvent { event, .. } => self.translator.translate(event),
|
||||
ProviderRunProjection::ModelRetry { .. }
|
||||
ProviderRunProjection::ModelTurnRequested { .. }
|
||||
| ProviderRunProjection::ModelTurnFinished { .. }
|
||||
| ProviderRunProjection::ModelRetry { .. }
|
||||
| ProviderRunProjection::ToolBatchReady { .. } => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -54,8 +54,12 @@ fn restored_provider_projection_skips_stream_initialization() {
|
||||
assert!(projector
|
||||
.project(ProviderRunProjection::ModelTurnStarted {
|
||||
work_id: work_id.clone(),
|
||||
profile: galaxy_agent_core::ProviderRequestProfile::new("base"),
|
||||
runtime_id: "runtime".to_owned(),
|
||||
model_id: "model".to_owned(),
|
||||
runtime_request_id: "request".to_owned(),
|
||||
retry_attempt: 0,
|
||||
elapsed_ms: 1,
|
||||
})
|
||||
.unwrap()
|
||||
.is_empty());
|
||||
|
||||
@@ -8,8 +8,9 @@ pub(crate) use event_translator::{
|
||||
ProviderRunResponseProjector, RuntimeResponseConfig, RuntimeResponseTranslator,
|
||||
};
|
||||
pub(crate) use provider_run_coordinator::{
|
||||
ProviderRunBlock, ProviderRunCoordinator, ProviderRunProfile, ProviderToolExecutionRef,
|
||||
ProviderToolLifecycleOutcome, BASE_PROVIDER_PROFILE, CLI_MONITOR_PROVIDER_PROFILE,
|
||||
ProviderRunBlock, ProviderRunCoordinator, ProviderRunProfile, ProviderRunProjection,
|
||||
ProviderToolExecutionRef, ProviderToolLifecycleOutcome, BASE_PROVIDER_PROFILE,
|
||||
CLI_MONITOR_PROVIDER_PROFILE,
|
||||
};
|
||||
pub(crate) use rig::{
|
||||
prepare_provider_run, provider_runtime_for_request, PreparedProviderRun, ProviderActionContext,
|
||||
|
||||
@@ -2,6 +2,7 @@ use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::error::Error;
|
||||
use std::fmt;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use futures::future::BoxFuture;
|
||||
use futures::StreamExt;
|
||||
@@ -12,18 +13,43 @@ use galaxy_agent_core::{
|
||||
ProviderRunOutcome, ProviderRunPhase, ProviderRunProtocolError, ProviderRunState,
|
||||
ProviderRunStep, RunEpoch, RuntimeKind, StopReason, ToolEvent, TurnControl, TurnRequest, Usage,
|
||||
};
|
||||
use instant::Instant;
|
||||
use warpui::r#async::FutureExt as _;
|
||||
|
||||
use crate::ai::agent::conversation::AIConversationId;
|
||||
|
||||
pub(crate) const BASE_PROVIDER_PROFILE: &str = "base";
|
||||
pub(crate) const CLI_MONITOR_PROVIDER_PROFILE: &str = "cli-monitor";
|
||||
const PROVIDER_MODEL_START_TIMEOUT: Duration = Duration::from_secs(120);
|
||||
const PROVIDER_MODEL_EVENT_IDLE_TIMEOUT: Duration = Duration::from_secs(300);
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub(crate) enum ProviderRunProjection {
|
||||
ModelTurnRequested {
|
||||
work_id: ExternalWorkId,
|
||||
profile: ProviderRequestProfile,
|
||||
runtime_id: String,
|
||||
model_id: String,
|
||||
retry_attempt: u32,
|
||||
},
|
||||
ModelTurnStarted {
|
||||
work_id: ExternalWorkId,
|
||||
profile: ProviderRequestProfile,
|
||||
runtime_id: String,
|
||||
model_id: String,
|
||||
runtime_request_id: String,
|
||||
retry_attempt: u32,
|
||||
elapsed_ms: u64,
|
||||
},
|
||||
ModelTurnFinished {
|
||||
work_id: ExternalWorkId,
|
||||
profile: ProviderRequestProfile,
|
||||
runtime_id: String,
|
||||
model_id: String,
|
||||
stop_reason: StopReason,
|
||||
retry_attempt: u32,
|
||||
elapsed_ms: u64,
|
||||
tool_call_count: usize,
|
||||
},
|
||||
ModelEvent {
|
||||
work_id: ExternalWorkId,
|
||||
@@ -31,7 +57,11 @@ pub(crate) enum ProviderRunProjection {
|
||||
},
|
||||
ModelRetry {
|
||||
work_id: ExternalWorkId,
|
||||
profile: ProviderRequestProfile,
|
||||
runtime_id: String,
|
||||
model_id: String,
|
||||
retry_attempt: u32,
|
||||
elapsed_ms: u64,
|
||||
error: AgentError,
|
||||
},
|
||||
ToolBatchReady {
|
||||
@@ -133,6 +163,8 @@ impl ProviderRunProfile {
|
||||
pub(crate) struct ProviderRunCoordinator {
|
||||
run: ProviderRun,
|
||||
profiles: BTreeMap<String, ProviderRunProfile>,
|
||||
model_start_timeout: Duration,
|
||||
model_event_idle_timeout: Duration,
|
||||
}
|
||||
|
||||
impl ProviderRunCoordinator {
|
||||
@@ -171,7 +203,18 @@ impl ProviderRunCoordinator {
|
||||
for (profile, config) in &profiles {
|
||||
validate_profile_runtime(profile, config.runtime.as_ref())?;
|
||||
}
|
||||
Ok(Self { run, profiles })
|
||||
Ok(Self {
|
||||
run,
|
||||
profiles,
|
||||
model_start_timeout: PROVIDER_MODEL_START_TIMEOUT,
|
||||
model_event_idle_timeout: PROVIDER_MODEL_EVENT_IDLE_TIMEOUT,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn set_model_timeouts(&mut self, start: Duration, event_idle: Duration) {
|
||||
self.model_start_timeout = start;
|
||||
self.model_event_idle_timeout = event_idle;
|
||||
}
|
||||
|
||||
pub(crate) fn run(&self) -> &ProviderRun {
|
||||
@@ -395,22 +438,67 @@ impl ProviderRunCoordinator {
|
||||
.iter()
|
||||
.map(|tool| tool.name.clone())
|
||||
.collect::<BTreeSet<_>>();
|
||||
let request = request_for_model_call(profile.request, &call);
|
||||
let stream = match profile.runtime.start_turn(request, control).await {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
self.handle_model_failure(&call.work_id, error, project)?;
|
||||
let runtime_id = profile.runtime.descriptor().id.clone();
|
||||
let model_id = profile.request.model.as_str().to_string();
|
||||
if !self.project_or_fail(
|
||||
ProviderRunProjection::ModelTurnRequested {
|
||||
work_id: call.work_id.clone(),
|
||||
profile: call.profile.clone(),
|
||||
runtime_id: runtime_id.clone(),
|
||||
model_id: model_id.clone(),
|
||||
retry_attempt: call.retry_attempt,
|
||||
},
|
||||
project,
|
||||
)? {
|
||||
return Ok(());
|
||||
}
|
||||
let request = request_for_model_call(profile.request.clone(), &call);
|
||||
let started_at = Instant::now();
|
||||
let stream = match profile
|
||||
.runtime
|
||||
.start_turn(request, control)
|
||||
.with_timeout(self.model_start_timeout)
|
||||
.await
|
||||
{
|
||||
Ok(Ok(stream)) => stream,
|
||||
Ok(Err(error)) => {
|
||||
self.handle_model_failure(&call, &profile, started_at, error, project)?;
|
||||
return Ok(());
|
||||
}
|
||||
Err(_) => {
|
||||
self.handle_model_failure(
|
||||
&call,
|
||||
&profile,
|
||||
started_at,
|
||||
provider_timeout_error("start", self.model_start_timeout),
|
||||
project,
|
||||
)?;
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
futures::pin_mut!(stream);
|
||||
let mut buffer = ModelTurnBuffer::default();
|
||||
|
||||
while let Some(event) = stream.next().await {
|
||||
let event = match event {
|
||||
Ok(event) => event,
|
||||
Err(error) => {
|
||||
self.handle_model_failure(&call.work_id, error, project)?;
|
||||
loop {
|
||||
let event = match stream
|
||||
.next()
|
||||
.with_timeout(self.model_event_idle_timeout)
|
||||
.await
|
||||
{
|
||||
Ok(Some(Ok(event))) => event,
|
||||
Ok(Some(Err(error))) => {
|
||||
self.handle_model_failure(&call, &profile, started_at, error, project)?;
|
||||
return Ok(());
|
||||
}
|
||||
Ok(None) => break,
|
||||
Err(_) => {
|
||||
self.handle_model_failure(
|
||||
&call,
|
||||
&profile,
|
||||
started_at,
|
||||
provider_timeout_error("event", self.model_event_idle_timeout),
|
||||
project,
|
||||
)?;
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
@@ -418,7 +506,9 @@ impl ProviderRunCoordinator {
|
||||
AgentEvent::TurnStarted { runtime_request_id } => {
|
||||
if buffer.started {
|
||||
self.handle_model_failure(
|
||||
&call.work_id,
|
||||
&call,
|
||||
&profile,
|
||||
started_at,
|
||||
protocol_error("provider emitted more than one TurnStarted event"),
|
||||
project,
|
||||
)?;
|
||||
@@ -426,7 +516,9 @@ impl ProviderRunCoordinator {
|
||||
}
|
||||
if runtime_request_id.is_empty() {
|
||||
self.handle_model_failure(
|
||||
&call.work_id,
|
||||
&call,
|
||||
&profile,
|
||||
started_at,
|
||||
protocol_error("provider emitted an empty runtime request ID"),
|
||||
project,
|
||||
)?;
|
||||
@@ -436,8 +528,12 @@ impl ProviderRunCoordinator {
|
||||
if !self.project_or_fail(
|
||||
ProviderRunProjection::ModelTurnStarted {
|
||||
work_id: call.work_id.clone(),
|
||||
profile: call.profile.clone(),
|
||||
runtime_id: runtime_id.clone(),
|
||||
model_id: model_id.clone(),
|
||||
runtime_request_id,
|
||||
retry_attempt: call.retry_attempt,
|
||||
elapsed_ms: elapsed_millis(started_at),
|
||||
},
|
||||
project,
|
||||
)? {
|
||||
@@ -445,7 +541,7 @@ impl ProviderRunCoordinator {
|
||||
}
|
||||
}
|
||||
AgentEvent::TextDelta { text } => {
|
||||
if !self.ensure_model_started(&call.work_id, &buffer, project)? {
|
||||
if !self.ensure_model_started(&call, &profile, started_at, &buffer, project)? {
|
||||
return Ok(());
|
||||
}
|
||||
buffer.text.push_str(&text);
|
||||
@@ -460,7 +556,7 @@ impl ProviderRunCoordinator {
|
||||
}
|
||||
}
|
||||
AgentEvent::ReasoningDelta { text } => {
|
||||
if !self.ensure_model_started(&call.work_id, &buffer, project)? {
|
||||
if !self.ensure_model_started(&call, &profile, started_at, &buffer, project)? {
|
||||
return Ok(());
|
||||
}
|
||||
buffer.reasoning.push_str(&text);
|
||||
@@ -475,7 +571,7 @@ impl ProviderRunCoordinator {
|
||||
}
|
||||
}
|
||||
AgentEvent::ReasoningCompleted { text, signature } => {
|
||||
if !self.ensure_model_started(&call.work_id, &buffer, project)? {
|
||||
if !self.ensure_model_started(&call, &profile, started_at, &buffer, project)? {
|
||||
return Ok(());
|
||||
}
|
||||
if !text.is_empty() {
|
||||
@@ -495,13 +591,13 @@ impl ProviderRunCoordinator {
|
||||
AgentEvent::Tool {
|
||||
event: ToolEvent::Proposed { call: tool_call },
|
||||
} => {
|
||||
if !self.ensure_model_started(&call.work_id, &buffer, project)? {
|
||||
if !self.ensure_model_started(&call, &profile, started_at, &buffer, project)? {
|
||||
return Ok(());
|
||||
}
|
||||
buffer.tool_calls.push(tool_call);
|
||||
}
|
||||
AgentEvent::UsageUpdated { usage } => {
|
||||
if !self.ensure_model_started(&call.work_id, &buffer, project)? {
|
||||
if !self.ensure_model_started(&call, &profile, started_at, &buffer, project)? {
|
||||
return Ok(());
|
||||
}
|
||||
buffer.usage.clone_from(&usage);
|
||||
@@ -519,7 +615,22 @@ impl ProviderRunCoordinator {
|
||||
}
|
||||
}
|
||||
AgentEvent::TurnStopped { reason } => {
|
||||
if !self.ensure_model_started(&call.work_id, &buffer, project)? {
|
||||
if !self.ensure_model_started(&call, &profile, started_at, &buffer, project)? {
|
||||
return Ok(());
|
||||
}
|
||||
if !self.project_or_fail(
|
||||
ProviderRunProjection::ModelTurnFinished {
|
||||
work_id: call.work_id.clone(),
|
||||
profile: call.profile.clone(),
|
||||
runtime_id: runtime_id.clone(),
|
||||
model_id: model_id.clone(),
|
||||
stop_reason: reason.clone(),
|
||||
retry_attempt: call.retry_attempt,
|
||||
elapsed_ms: elapsed_millis(started_at),
|
||||
tool_call_count: buffer.tool_calls.len(),
|
||||
},
|
||||
project,
|
||||
)? {
|
||||
return Ok(());
|
||||
}
|
||||
if reason == StopReason::Cancelled {
|
||||
@@ -547,7 +658,9 @@ impl ProviderRunCoordinator {
|
||||
| AgentEvent::UserInputAccepted { .. }
|
||||
| AgentEvent::RuntimeNotice { .. } => {
|
||||
self.handle_model_failure(
|
||||
&call.work_id,
|
||||
&call,
|
||||
&profile,
|
||||
started_at,
|
||||
protocol_error(
|
||||
"direct-provider transport emitted a non-model lifecycle event",
|
||||
),
|
||||
@@ -559,7 +672,9 @@ impl ProviderRunCoordinator {
|
||||
}
|
||||
|
||||
self.handle_model_failure(
|
||||
&call.work_id,
|
||||
&call,
|
||||
&profile,
|
||||
started_at,
|
||||
protocol_error("provider stream ended before TurnStopped"),
|
||||
project,
|
||||
)?;
|
||||
@@ -568,7 +683,9 @@ impl ProviderRunCoordinator {
|
||||
|
||||
fn ensure_model_started<F>(
|
||||
&mut self,
|
||||
work_id: &ExternalWorkId,
|
||||
call: &ProviderModelCall,
|
||||
profile: &ProviderRunProfile,
|
||||
started_at: Instant,
|
||||
buffer: &ModelTurnBuffer,
|
||||
project: &mut F,
|
||||
) -> Result<bool, ProviderRunCoordinatorError>
|
||||
@@ -579,7 +696,9 @@ impl ProviderRunCoordinator {
|
||||
return Ok(true);
|
||||
}
|
||||
self.handle_model_failure(
|
||||
work_id,
|
||||
call,
|
||||
profile,
|
||||
started_at,
|
||||
protocol_error("provider emitted model output before TurnStarted"),
|
||||
project,
|
||||
)?;
|
||||
@@ -588,14 +707,18 @@ impl ProviderRunCoordinator {
|
||||
|
||||
fn handle_model_failure<F>(
|
||||
&mut self,
|
||||
work_id: &ExternalWorkId,
|
||||
call: &ProviderModelCall,
|
||||
profile: &ProviderRunProfile,
|
||||
started_at: Instant,
|
||||
error: AgentError,
|
||||
project: &mut F,
|
||||
) -> Result<(), ProviderRunCoordinatorError>
|
||||
where
|
||||
F: FnMut(ProviderRunProjection) -> Result<(), String>,
|
||||
{
|
||||
let disposition = self.run.register_model_failure(work_id, error.clone())?;
|
||||
let disposition = self
|
||||
.run
|
||||
.register_model_failure(&call.work_id, error.clone())?;
|
||||
if disposition == ModelFailureDisposition::RetryScheduled {
|
||||
let retry_attempt = match self.run.state() {
|
||||
ProviderRunState::AwaitingModel { call } => call.retry_attempt,
|
||||
@@ -616,8 +739,12 @@ impl ProviderRunCoordinator {
|
||||
};
|
||||
self.project_or_fail(
|
||||
ProviderRunProjection::ModelRetry {
|
||||
work_id: work_id.clone(),
|
||||
work_id: call.work_id.clone(),
|
||||
profile: call.profile.clone(),
|
||||
runtime_id: profile.runtime.descriptor().id.clone(),
|
||||
model_id: profile.request.model.as_str().to_string(),
|
||||
retry_attempt,
|
||||
elapsed_ms: elapsed_millis(started_at),
|
||||
error,
|
||||
},
|
||||
project,
|
||||
@@ -683,6 +810,22 @@ impl ModelTurnBuffer {
|
||||
}
|
||||
}
|
||||
|
||||
fn elapsed_millis(started_at: Instant) -> u64 {
|
||||
u64::try_from(started_at.elapsed().as_millis()).unwrap_or(u64::MAX)
|
||||
}
|
||||
|
||||
fn provider_timeout_error(stage: &str, timeout: Duration) -> AgentError {
|
||||
let mut error = AgentError::new(
|
||||
AgentErrorKind::Transport,
|
||||
format!(
|
||||
"provider model {stage} timed out after {} seconds",
|
||||
timeout.as_secs()
|
||||
),
|
||||
);
|
||||
error.recoverable = true;
|
||||
error
|
||||
}
|
||||
|
||||
fn validate_profile_runtime(
|
||||
profile: &str,
|
||||
runtime: &dyn AgentRuntime,
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use std::collections::VecDeque;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::Mutex;
|
||||
use std::time::Duration;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use galaxy_agent_core::{
|
||||
@@ -66,6 +68,65 @@ impl AgentRuntime for ScriptedRuntime {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum FirstAttemptStall {
|
||||
Start,
|
||||
Event,
|
||||
}
|
||||
|
||||
struct StallingRuntime {
|
||||
descriptor: RuntimeDescriptor,
|
||||
first_attempt_stall: FirstAttemptStall,
|
||||
attempts: AtomicUsize,
|
||||
requests: Mutex<Vec<TurnRequest>>,
|
||||
}
|
||||
|
||||
impl StallingRuntime {
|
||||
fn new(first_attempt_stall: FirstAttemptStall) -> Self {
|
||||
Self {
|
||||
descriptor: RuntimeDescriptor {
|
||||
id: "stalling".to_string(),
|
||||
display_name: "Stalling provider".to_string(),
|
||||
kind: RuntimeKind::Provider,
|
||||
capabilities: RuntimeCapabilities::provider(),
|
||||
},
|
||||
first_attempt_stall,
|
||||
attempts: AtomicUsize::new(0),
|
||||
requests: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn requests(&self) -> Vec<TurnRequest> {
|
||||
self.requests.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl AgentRuntime for StallingRuntime {
|
||||
fn descriptor(&self) -> &RuntimeDescriptor {
|
||||
&self.descriptor
|
||||
}
|
||||
|
||||
async fn start_turn(
|
||||
&self,
|
||||
request: TurnRequest,
|
||||
_control: TurnControl,
|
||||
) -> Result<AgentEventStream, AgentError> {
|
||||
self.requests.lock().unwrap().push(request);
|
||||
let attempt = self.attempts.fetch_add(1, Ordering::Relaxed);
|
||||
if attempt == 0 {
|
||||
match self.first_attempt_stall {
|
||||
FirstAttemptStall::Start => return futures::future::pending().await,
|
||||
FirstAttemptStall::Event => {
|
||||
let started = futures::stream::iter(vec![started("request-stalled")]);
|
||||
return Ok(Box::pin(started.chain(futures::stream::pending())));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(Box::pin(futures::stream::iter(answer_turn().unwrap())))
|
||||
}
|
||||
}
|
||||
|
||||
fn request() -> TurnRequest {
|
||||
let mut request = TurnRequest::new(
|
||||
"test-model",
|
||||
@@ -690,7 +751,9 @@ async fn recoverable_start_failure_retries_the_same_work_identity() {
|
||||
retry_attempt,
|
||||
..
|
||||
} => Some((work_id.clone(), *retry_attempt)),
|
||||
ProviderRunProjection::ModelTurnStarted { .. }
|
||||
ProviderRunProjection::ModelTurnRequested { .. }
|
||||
| ProviderRunProjection::ModelTurnStarted { .. }
|
||||
| ProviderRunProjection::ModelTurnFinished { .. }
|
||||
| ProviderRunProjection::ModelEvent { .. }
|
||||
| ProviderRunProjection::ToolBatchReady { .. } => None,
|
||||
})
|
||||
@@ -703,7 +766,9 @@ async fn recoverable_start_failure_retries_the_same_work_identity() {
|
||||
retry_attempt: 1,
|
||||
..
|
||||
} => Some(work_id.clone()),
|
||||
ProviderRunProjection::ModelTurnStarted { .. }
|
||||
ProviderRunProjection::ModelTurnRequested { .. }
|
||||
| ProviderRunProjection::ModelTurnStarted { .. }
|
||||
| ProviderRunProjection::ModelTurnFinished { .. }
|
||||
| ProviderRunProjection::ModelEvent { .. }
|
||||
| ProviderRunProjection::ModelRetry { .. }
|
||||
| ProviderRunProjection::ToolBatchReady { .. } => None,
|
||||
@@ -713,6 +778,119 @@ async fn recoverable_start_failure_retries_the_same_work_identity() {
|
||||
assert_eq!(retry.1, 1);
|
||||
}
|
||||
|
||||
fn assert_single_retry_lifecycle(
|
||||
projections: &[ProviderRunProjection],
|
||||
expected_timeout_stage: &str,
|
||||
expected_initial_start: bool,
|
||||
) {
|
||||
let mut work_ids = Vec::new();
|
||||
let mut phases = Vec::new();
|
||||
let mut retry_error = None;
|
||||
for projection in projections {
|
||||
match projection {
|
||||
ProviderRunProjection::ModelTurnRequested {
|
||||
work_id,
|
||||
retry_attempt,
|
||||
..
|
||||
} => {
|
||||
work_ids.push(work_id.clone());
|
||||
phases.push(format!("requested:{retry_attempt}"));
|
||||
}
|
||||
ProviderRunProjection::ModelTurnStarted {
|
||||
work_id,
|
||||
retry_attempt,
|
||||
..
|
||||
} => {
|
||||
work_ids.push(work_id.clone());
|
||||
phases.push(format!("started:{retry_attempt}"));
|
||||
}
|
||||
ProviderRunProjection::ModelRetry {
|
||||
work_id,
|
||||
retry_attempt,
|
||||
error,
|
||||
..
|
||||
} => {
|
||||
work_ids.push(work_id.clone());
|
||||
phases.push(format!("retry:{retry_attempt}"));
|
||||
retry_error = Some(error);
|
||||
}
|
||||
ProviderRunProjection::ModelTurnFinished {
|
||||
work_id,
|
||||
retry_attempt,
|
||||
..
|
||||
} => {
|
||||
work_ids.push(work_id.clone());
|
||||
phases.push(format!("finished:{retry_attempt}"));
|
||||
}
|
||||
ProviderRunProjection::ModelEvent { .. }
|
||||
| ProviderRunProjection::ToolBatchReady { .. } => {}
|
||||
}
|
||||
}
|
||||
|
||||
let expected = if expected_initial_start {
|
||||
vec![
|
||||
"requested:0",
|
||||
"started:0",
|
||||
"retry:1",
|
||||
"requested:1",
|
||||
"started:1",
|
||||
"finished:1",
|
||||
]
|
||||
} else {
|
||||
vec![
|
||||
"requested:0",
|
||||
"retry:1",
|
||||
"requested:1",
|
||||
"started:1",
|
||||
"finished:1",
|
||||
]
|
||||
};
|
||||
assert_eq!(phases, expected);
|
||||
assert!(work_ids.windows(2).all(|ids| ids[0] == ids[1]));
|
||||
let retry_error = retry_error.expect("timeout retry error");
|
||||
assert_eq!(retry_error.kind, AgentErrorKind::Transport);
|
||||
assert!(retry_error.recoverable);
|
||||
assert!(retry_error.message.contains(expected_timeout_stage));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn model_start_timeout_retries_the_same_work_identity() {
|
||||
let runtime = Arc::new(StallingRuntime::new(FirstAttemptStall::Start));
|
||||
let mut coordinator = coordinator(runtime.clone());
|
||||
coordinator.set_model_timeouts(Duration::from_millis(10), Duration::from_secs(1));
|
||||
let mut projections = Vec::new();
|
||||
let (_sender, control) = turn_control();
|
||||
|
||||
let block = coordinator
|
||||
.drive_until_blocked(control, collect_projection(&mut projections))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(matches!(block, ProviderRunBlock::AwaitingDriver { .. }));
|
||||
assert_eq!(runtime.requests().len(), 2);
|
||||
assert_eq!(coordinator.run().model_retries(), 1);
|
||||
assert_single_retry_lifecycle(&projections, "start timed out", false);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn model_event_idle_timeout_retries_the_same_work_identity() {
|
||||
let runtime = Arc::new(StallingRuntime::new(FirstAttemptStall::Event));
|
||||
let mut coordinator = coordinator(runtime.clone());
|
||||
coordinator.set_model_timeouts(Duration::from_secs(1), Duration::from_millis(10));
|
||||
let mut projections = Vec::new();
|
||||
let (_sender, control) = turn_control();
|
||||
|
||||
let block = coordinator
|
||||
.drive_until_blocked(control, collect_projection(&mut projections))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(matches!(block, ProviderRunBlock::AwaitingDriver { .. }));
|
||||
assert_eq!(runtime.requests().len(), 2);
|
||||
assert_eq!(coordinator.run().model_retries(), 1);
|
||||
assert_single_retry_lifecycle(&projections, "event timed out", true);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tool_projection_failure_synthesizes_a_result_and_fails_the_run() {
|
||||
let runtime = Arc::new(ScriptedRuntime::new(vec![tool_turn()]));
|
||||
@@ -724,7 +902,9 @@ async fn tool_projection_failure_synthesizes_a_result_and_fails_the_run() {
|
||||
ProviderRunProjection::ToolBatchReady { .. } => {
|
||||
Err("task projection disappeared".to_string())
|
||||
}
|
||||
ProviderRunProjection::ModelTurnStarted { .. }
|
||||
ProviderRunProjection::ModelTurnRequested { .. }
|
||||
| ProviderRunProjection::ModelTurnStarted { .. }
|
||||
| ProviderRunProjection::ModelTurnFinished { .. }
|
||||
| ProviderRunProjection::ModelEvent { .. }
|
||||
| ProviderRunProjection::ModelRetry { .. } => Ok(()),
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user