Improve provider reliability and usage visibility
This commit is contained in:
@@ -40,6 +40,7 @@ pub(crate) struct RuntimeResponseTranslator {
|
||||
activity_message_ids: HashMap<String, String>,
|
||||
activities: HashMap<String, RuntimeActivity>,
|
||||
has_visible_output: bool,
|
||||
/// Usage for the most recent model call.
|
||||
usage: Usage,
|
||||
context_usage: Option<(u64, u64)>,
|
||||
}
|
||||
@@ -107,25 +108,22 @@ impl ProviderRunResponseProjector {
|
||||
pub(crate) fn finish(
|
||||
&mut self,
|
||||
outcome: &ProviderRunOutcome,
|
||||
aggregate_usage: &Usage,
|
||||
) -> Result<Vec<ResponseEvent>, String> {
|
||||
if self.finished {
|
||||
return Err("provider run projection is already finished".to_string());
|
||||
}
|
||||
self.finished = true;
|
||||
match outcome {
|
||||
ProviderRunOutcome::Completed(completion) => {
|
||||
self.translator.translate(AgentEvent::TurnStopped {
|
||||
reason: completion.stop_reason.clone(),
|
||||
})
|
||||
}
|
||||
ProviderRunOutcome::Failed(failure) => {
|
||||
Ok(self.translator.provider_failure(&failure.message))
|
||||
}
|
||||
ProviderRunOutcome::Cancelled { .. } => {
|
||||
self.translator.translate(AgentEvent::TurnStopped {
|
||||
reason: StopReason::Cancelled,
|
||||
})
|
||||
}
|
||||
ProviderRunOutcome::Completed(completion) => Ok(self
|
||||
.translator
|
||||
.finish_provider_run(completion.stop_reason.clone(), aggregate_usage)),
|
||||
ProviderRunOutcome::Failed(failure) => Ok(self
|
||||
.translator
|
||||
.provider_failure(&failure.message, aggregate_usage)),
|
||||
ProviderRunOutcome::Cancelled { .. } => Ok(self
|
||||
.translator
|
||||
.finish_provider_run(StopReason::Cancelled, aggregate_usage)),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -165,6 +163,7 @@ impl RuntimeResponseTranslator {
|
||||
let mut events = Vec::new();
|
||||
match event {
|
||||
AgentEvent::TurnStarted { .. } => self.initialize(&mut events),
|
||||
AgentEvent::KeepAlive => {}
|
||||
AgentEvent::TextDelta { text } => {
|
||||
self.initialize(&mut events);
|
||||
self.add_or_append_text(&text, &mut events);
|
||||
@@ -249,13 +248,31 @@ impl RuntimeResponseTranslator {
|
||||
events
|
||||
}
|
||||
|
||||
fn finish_provider_run(
|
||||
&mut self,
|
||||
reason: StopReason,
|
||||
aggregate_usage: &Usage,
|
||||
) -> Vec<ResponseEvent> {
|
||||
let mut events = Vec::new();
|
||||
self.initialize(&mut events);
|
||||
if !self.has_visible_output && reason != StopReason::Cancelled {
|
||||
if let Some(message) = self.config.empty_output_message.clone() {
|
||||
self.add_or_append_text(&message, &mut events);
|
||||
}
|
||||
}
|
||||
events.push(self.finished_with_usage(map_stop_reason(reason), aggregate_usage));
|
||||
events
|
||||
}
|
||||
|
||||
pub(crate) fn begin_followup_turn(&mut self) {
|
||||
self.text_message_id = None;
|
||||
self.reasoning_message_id = None;
|
||||
self.usage = Usage::default();
|
||||
}
|
||||
|
||||
fn discard_failed_turn_output(&mut self) -> Vec<ResponseEvent> {
|
||||
let mut events = Vec::new();
|
||||
self.usage = Usage::default();
|
||||
if let Some(message_id) = self.text_message_id.take() {
|
||||
events.push(build_replace_text_message(
|
||||
&self.config.task_id,
|
||||
@@ -404,20 +421,27 @@ impl RuntimeResponseTranslator {
|
||||
self.finished_with_reason(map_stop_reason(reason))
|
||||
}
|
||||
|
||||
fn provider_failure(&mut self, message: &str) -> Vec<ResponseEvent> {
|
||||
fn provider_failure(&mut self, message: &str, aggregate_usage: &Usage) -> Vec<ResponseEvent> {
|
||||
let mut events = Vec::new();
|
||||
self.initialize(&mut events);
|
||||
events.push(
|
||||
self.finished_with_reason(stream_finished::Reason::InternalError(
|
||||
stream_finished::InternalError {
|
||||
message: message.to_owned(),
|
||||
},
|
||||
)),
|
||||
);
|
||||
events.push(self.finished_with_usage(
|
||||
stream_finished::Reason::InternalError(stream_finished::InternalError {
|
||||
message: message.to_owned(),
|
||||
}),
|
||||
aggregate_usage,
|
||||
));
|
||||
events
|
||||
}
|
||||
|
||||
fn finished_with_reason(&self, reason: stream_finished::Reason) -> ResponseEvent {
|
||||
self.finished_with_usage(reason, &self.usage)
|
||||
}
|
||||
|
||||
fn finished_with_usage(
|
||||
&self,
|
||||
reason: stream_finished::Reason,
|
||||
aggregate_usage: &Usage,
|
||||
) -> ResponseEvent {
|
||||
if !self.config.capabilities.host_managed_history {
|
||||
let (used_tokens, context_size) = self.context_usage.unwrap_or_default();
|
||||
return build_context_finished(
|
||||
@@ -430,10 +454,16 @@ impl RuntimeResponseTranslator {
|
||||
build_stream_finished(
|
||||
reason,
|
||||
StreamUsage {
|
||||
input_tokens: saturating_i32(self.usage.input_tokens),
|
||||
output_tokens: saturating_i32(self.usage.output_tokens),
|
||||
cache_read_tokens: saturating_i32(self.usage.cached_input_tokens),
|
||||
cache_write_tokens: saturating_i32(self.usage.cache_creation_input_tokens),
|
||||
input_tokens: saturating_i32(aggregate_usage.input_tokens),
|
||||
output_tokens: saturating_i32(aggregate_usage.output_tokens),
|
||||
cache_read_tokens: saturating_i32(aggregate_usage.cached_input_tokens),
|
||||
cache_write_tokens: saturating_i32(aggregate_usage.cache_creation_input_tokens),
|
||||
current_context_tokens: Some(saturating_i32(
|
||||
self.usage
|
||||
.input_tokens
|
||||
.saturating_add(self.usage.cached_input_tokens)
|
||||
.saturating_add(self.usage.cache_creation_input_tokens),
|
||||
)),
|
||||
cost_in_cents: 0.0,
|
||||
model_id: self.config.model_id.clone(),
|
||||
max_context_tokens: self.config.max_context_tokens,
|
||||
|
||||
@@ -404,6 +404,7 @@ fn provider_retry_clears_failed_attempt_output_before_new_messages() {
|
||||
runtime_id: "runtime".to_owned(),
|
||||
model_id: "model".to_owned(),
|
||||
retry_attempt: 1,
|
||||
max_retries: 3,
|
||||
elapsed_ms: 2,
|
||||
error: galaxy_agent_core::AgentError::new(
|
||||
galaxy_agent_core::AgentErrorKind::Transport,
|
||||
|
||||
@@ -67,6 +67,7 @@ pub(crate) enum ProviderRunProjection {
|
||||
runtime_id: String,
|
||||
model_id: String,
|
||||
retry_attempt: u32,
|
||||
max_retries: u32,
|
||||
elapsed_ms: u64,
|
||||
error: AgentError,
|
||||
},
|
||||
@@ -639,6 +640,7 @@ impl ProviderRunCoordinator {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
AgentEvent::KeepAlive => {}
|
||||
AgentEvent::TextDelta { text } => {
|
||||
if !self
|
||||
.ensure_model_started_acknowledged(
|
||||
@@ -732,14 +734,11 @@ impl ProviderRunCoordinator {
|
||||
return Ok(());
|
||||
}
|
||||
buffer.usage.clone_from(&usage);
|
||||
let cumulative_usage = combined_usage(self.run.usage(), &usage);
|
||||
if !self
|
||||
.project_or_fail_acknowledged(
|
||||
ProviderRunProjection::ModelEvent {
|
||||
work_id: call.work_id.clone(),
|
||||
event: AgentEvent::UsageUpdated {
|
||||
usage: cumulative_usage,
|
||||
},
|
||||
event: AgentEvent::UsageUpdated { usage },
|
||||
},
|
||||
project,
|
||||
)
|
||||
@@ -890,6 +889,7 @@ impl ProviderRunCoordinator {
|
||||
runtime_id: profile.runtime.descriptor().id.clone(),
|
||||
model_id: profile.request.model.as_str().to_string(),
|
||||
retry_attempt,
|
||||
max_retries: self.run.max_model_retries_per_turn(),
|
||||
elapsed_ms: elapsed_millis(started_at),
|
||||
error,
|
||||
},
|
||||
@@ -1018,19 +1018,6 @@ fn request_for_model_call(mut template: TurnRequest, call: &ProviderModelCall) -
|
||||
template
|
||||
}
|
||||
|
||||
fn combined_usage(previous: &Usage, current: &Usage) -> Usage {
|
||||
Usage {
|
||||
input_tokens: previous.input_tokens.saturating_add(current.input_tokens),
|
||||
output_tokens: previous.output_tokens.saturating_add(current.output_tokens),
|
||||
cached_input_tokens: previous
|
||||
.cached_input_tokens
|
||||
.saturating_add(current.cached_input_tokens),
|
||||
cache_creation_input_tokens: previous
|
||||
.cache_creation_input_tokens
|
||||
.saturating_add(current.cache_creation_input_tokens),
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_event_call_id(event: &ToolEvent) -> Result<&str, ProviderRunCoordinatorError> {
|
||||
match event {
|
||||
ToolEvent::Proposed { call } => Ok(&call.id),
|
||||
|
||||
@@ -363,7 +363,7 @@ async fn one_run_drives_model_tool_and_followup_turns_with_atomic_history() {
|
||||
ProviderRunProjection::ModelEvent {
|
||||
event: AgentEvent::UsageUpdated { usage },
|
||||
..
|
||||
} if usage.input_tokens == 30 && usage.output_tokens == 9
|
||||
} if usage.input_tokens == 20 && usage.output_tokens == 5
|
||||
)));
|
||||
|
||||
coordinator.run_mut().complete(&work_id).unwrap();
|
||||
@@ -761,8 +761,9 @@ async fn recoverable_start_failure_retries_the_same_work_identity() {
|
||||
ProviderRunProjection::ModelRetry {
|
||||
work_id,
|
||||
retry_attempt,
|
||||
max_retries,
|
||||
..
|
||||
} => Some((work_id.clone(), *retry_attempt)),
|
||||
} => Some((work_id.clone(), *retry_attempt, *max_retries)),
|
||||
ProviderRunProjection::ModelTurnRequested { .. }
|
||||
| ProviderRunProjection::ModelTurnStarted { .. }
|
||||
| ProviderRunProjection::ModelTurnFinished { .. }
|
||||
@@ -788,6 +789,7 @@ async fn recoverable_start_failure_retries_the_same_work_identity() {
|
||||
.expect("retried model start");
|
||||
assert_eq!(retry.0, started);
|
||||
assert_eq!(retry.1, 1);
|
||||
assert_eq!(retry.2, 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -969,6 +971,38 @@ async fn model_event_idle_timeout_retries_the_same_work_identity() {
|
||||
assert_single_retry_lifecycle(&projections, "event timed out", true);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn model_keepalive_preserves_the_active_turn_without_rendering_output() {
|
||||
let runtime = Arc::new(ScriptedRuntime::new(vec![Ok(vec![
|
||||
started("request-keepalive"),
|
||||
Ok(AgentEvent::KeepAlive),
|
||||
Ok(AgentEvent::TextDelta {
|
||||
text: "finished".to_string(),
|
||||
}),
|
||||
stopped(StopReason::Completed),
|
||||
])]));
|
||||
let mut coordinator = coordinator(runtime);
|
||||
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!(coordinator.run().model_retries(), 0);
|
||||
assert!(!projections.iter().any(|projection| {
|
||||
matches!(
|
||||
projection,
|
||||
ProviderRunProjection::ModelEvent {
|
||||
event: AgentEvent::KeepAlive,
|
||||
..
|
||||
}
|
||||
)
|
||||
}));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn persistent_checkpoint_failure_terminates_without_redrive() {
|
||||
let runtime = Arc::new(ScriptedRuntime::new(vec![answer_turn()]));
|
||||
@@ -1168,7 +1202,11 @@ async fn transcript_projector_emits_one_ui_stream_for_the_whole_run() {
|
||||
else {
|
||||
panic!("expected terminal run");
|
||||
};
|
||||
ui_events.extend(projector.finish(&outcome).unwrap());
|
||||
ui_events.extend(
|
||||
projector
|
||||
.finish(&outcome, coordinator.run().usage())
|
||||
.unwrap(),
|
||||
);
|
||||
assert_eq!(
|
||||
count_response_events(&ui_events, ResponseEventKind::Finished),
|
||||
1
|
||||
@@ -1188,11 +1226,14 @@ fn transcript_projector_preserves_provider_failure_message() {
|
||||
empty_output_message: None,
|
||||
});
|
||||
let events = projector
|
||||
.finish(&ProviderRunOutcome::Failed(ProviderRunFailure {
|
||||
kind: ProviderRunFailureKind::ModelCall,
|
||||
message: "upstream provider rejected the request".to_string(),
|
||||
source: None,
|
||||
}))
|
||||
.finish(
|
||||
&ProviderRunOutcome::Failed(ProviderRunFailure {
|
||||
kind: ProviderRunFailureKind::ModelCall,
|
||||
message: "upstream provider rejected the request".to_string(),
|
||||
source: None,
|
||||
}),
|
||||
&Usage::default(),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let finished = events
|
||||
|
||||
@@ -651,6 +651,7 @@ fn build_system_prompt(
|
||||
);
|
||||
let contexts = inputs.iter().filter_map(AIAgentInput::context).flatten();
|
||||
let mut environment = Vec::new();
|
||||
let mut request_time = None;
|
||||
let mut project_rules = Vec::new();
|
||||
let mut available_skills = Vec::new();
|
||||
let mut attached_context = Vec::new();
|
||||
@@ -724,7 +725,7 @@ fn build_system_prompt(
|
||||
attached_context.push(("Selected text".to_string(), text.clone()));
|
||||
}
|
||||
AIAgentContext::CurrentTime { current_time } => {
|
||||
environment.push(format!("Current time: {current_time}"));
|
||||
request_time = Some(*current_time);
|
||||
}
|
||||
AIAgentContext::Codebase { path, name } => {
|
||||
environment.push(format!("Indexed codebase: {name} ({path})"));
|
||||
@@ -851,6 +852,14 @@ fn build_system_prompt(
|
||||
);
|
||||
}
|
||||
}
|
||||
// Keep volatile request data at the end of the system prompt. Provider
|
||||
// prompt caches match the longest exact prefix, so putting the current
|
||||
// timestamp ahead of rules and tool instructions invalidates that stable
|
||||
// prefix on every model call.
|
||||
if let Some(request_time) = request_time {
|
||||
prompt.push_str("\n## Request Time\n");
|
||||
prompt.push_str(&format!("- Current time: {request_time}\n"));
|
||||
}
|
||||
prompt
|
||||
}
|
||||
|
||||
|
||||
@@ -81,6 +81,32 @@ fn native_context_reaches_rig_without_a_proto_context_conversion() {
|
||||
assert!(prompt.contains("Indexed codebase: galaxy (/repo)"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn volatile_request_time_follows_the_cacheable_system_prompt_prefix() {
|
||||
let mut params = RequestParams::new_for_test();
|
||||
params.global_rules = vec![(
|
||||
"Stable rule".to_string(),
|
||||
"Preserve this cacheable instruction.".to_string(),
|
||||
)];
|
||||
params.input = vec![user_query_with_context(
|
||||
"Inspect the cache layout",
|
||||
vec![AIAgentContext::CurrentTime {
|
||||
current_time: chrono::Local::now(),
|
||||
}],
|
||||
)];
|
||||
|
||||
let prepared = prepare_rig_turn(&config(), params, vec![ToolType::ReadFiles], Vec::new());
|
||||
let prompt = prepared.request.system_prompt.expect("system prompt");
|
||||
let rule_position = prompt
|
||||
.find("Preserve this cacheable instruction.")
|
||||
.expect("global rule");
|
||||
let tools_position = prompt.find("## Available Tools").expect("tool contract");
|
||||
let time_position = prompt.find("## Request Time").expect("request time");
|
||||
|
||||
assert!(rule_position < time_position);
|
||||
assert!(tools_position < time_position);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_a_rig_turn_directly_from_galaxy_request_state() {
|
||||
let mut params = RequestParams::new_for_test();
|
||||
|
||||
Reference in New Issue
Block a user