diff --git a/crates/merry-cli/src/debug/openai/tests.rs b/crates/merry-cli/src/debug/openai/tests.rs index c08b652e..01b5b168 100644 --- a/crates/merry-cli/src/debug/openai/tests.rs +++ b/crates/merry-cli/src/debug/openai/tests.rs @@ -68,10 +68,12 @@ async fn tool_helper_executes_one_pending_call_and_continues() { [ "session_started", "step_started", + "model_output_rate_updated", "tool_call_pending", "artifact_recorded", "tool_call_resolved", "step_started", + "model_output_rate_updated", "assistant_output_recorded", "step_completed", ] @@ -259,6 +261,7 @@ async fn tool_helper_errors_when_first_step_does_not_call_debug_echo() { [ "session_started", "step_started", + "model_output_rate_updated", "assistant_output_recorded", "step_completed", ] diff --git a/crates/merry-cli/src/runtime_events.rs b/crates/merry-cli/src/runtime_events.rs index 7cdca1c7..67597fc1 100644 --- a/crates/merry-cli/src/runtime_events.rs +++ b/crates/merry-cli/src/runtime_events.rs @@ -118,7 +118,7 @@ mod tests { let text = String::from_utf8(output).expect("output should be utf-8"); let lines = text.lines().collect::>(); - assert_eq!(lines.len(), 5); + assert_eq!(lines.len(), 6); let event_types = lines .iter() .map(|line| { @@ -132,6 +132,7 @@ mod tests { [ "session_started", "step_started", + "model_output_rate_updated", "assistant_output_delta", "assistant_output_recorded", "step_completed" diff --git a/crates/merry-cli/src/tui/controller.rs b/crates/merry-cli/src/tui/controller.rs index f9486ba1..53260c74 100644 --- a/crates/merry-cli/src/tui/controller.rs +++ b/crates/merry-cli/src/tui/controller.rs @@ -277,8 +277,13 @@ pub(crate) async fn run_controller( }; match message { InteractiveRunMessage::Event(event) => { + let stream_progress = matches!(event, + merry_core::RuntimeEvent::AssistantMessageDelta { .. } + | merry_core::RuntimeEvent::ModelOutputRateUpdated { .. }); projector.apply(event, &mut state); - render_once(&mut terminal, &mut state)?; + if !stream_progress || !state.is_active_run() { + render_once(&mut terminal, &mut state)?; + } } InteractiveRunMessage::ToolInvocations { batch } => { return Err(unexpected(format!( diff --git a/crates/merry-cli/src/tui/mod.rs b/crates/merry-cli/src/tui/mod.rs index 7163695c..2d18756f 100644 --- a/crates/merry-cli/src/tui/mod.rs +++ b/crates/merry-cli/src/tui/mod.rs @@ -39,6 +39,7 @@ mod input_history_store; pub(crate) mod keymap; mod layout; mod markdown; +mod output_rate; mod overlay; mod overlay_render; mod plan; diff --git a/crates/merry-cli/src/tui/output_rate.rs b/crates/merry-cli/src/tui/output_rate.rs new file mode 100644 index 00000000..120c8471 --- /dev/null +++ b/crates/merry-cli/src/tui/output_rate.rs @@ -0,0 +1,34 @@ +use merry_core::{ModelOutputRate, RuntimeEvent}; + +/// Presentation of the last measurable runtime-owned throughput sample. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct OutputRate { + last_measurable_sample: Option, +} + +impl OutputRate { + pub(crate) fn observe(&mut self, event: &RuntimeEvent) { + let RuntimeEvent::ModelOutputRateUpdated { + rate: Some(rate), .. + } = event + else { + return; + }; + if rate.tokens_per_second().is_some() || self.last_measurable_sample.is_none() { + self.last_measurable_sample = Some(*rate); + } + } + + pub(crate) fn label(&self) -> String { + match self + .last_measurable_sample + .and_then(|sample| sample.tokens_per_second().map(|rate| (sample, rate))) + { + Some((sample, rate)) => { + let prefix = if sample.is_estimated() { "≈" } else { "" }; + format!("{prefix}{rate:.1} tok/s") + } + None => "- tok/s".to_owned(), + } + } +} diff --git a/crates/merry-cli/src/tui/projector.rs b/crates/merry-cli/src/tui/projector.rs index 8328b5a4..f36b6177 100644 --- a/crates/merry-cli/src/tui/projector.rs +++ b/crates/merry-cli/src/tui/projector.rs @@ -83,6 +83,7 @@ impl TuiProjector { } pub(crate) fn apply(&mut self, event: RuntimeEvent, state: &mut TuiState) { + state.observe_output_rate(&event); match event { RuntimeEvent::AssistantMessage { text, .. } => { if let Some(index) = self.streaming_assistant_index.take() { diff --git a/crates/merry-cli/src/tui/state.rs b/crates/merry-cli/src/tui/state.rs index 3e1bc1c6..426c7e4c 100644 --- a/crates/merry-cli/src/tui/state.rs +++ b/crates/merry-cli/src/tui/state.rs @@ -2,6 +2,7 @@ use crate::tui::{ completion::{CompletionMenu, CompletionSources}, input::{DraftImage, InputHistory, TextInput, TuiSubmission}, keymap::Keymap, + output_rate::OutputRate, plan::PlanUiState, preferences::{TuiPreferences, TuiSettingsDefaults}, status::{format_header_status_parts, format_session_usage_full}, @@ -9,7 +10,7 @@ use crate::tui::{ theme::TuiTheme, }; use clipboard::ClipboardFeedback; -use merry_core::{InteractiveRunState, QueuedInputLane, SessionUsage}; +use merry_core::{InteractiveRunState, QueuedInputLane, RuntimeEvent, SessionUsage}; use merry_runtime::SkillMetadata; use overlays::OverlayState; use std::{ @@ -61,6 +62,7 @@ pub(crate) struct TuiState { last_completed_run_elapsed: Option, pending_empty_input_quit: bool, usage: Option, + output_rate: OutputRate, overlays: OverlayState, preferences: TuiPreferences, settings_defaults: TuiSettingsDefaults, @@ -119,6 +121,7 @@ impl TuiState { last_completed_run_elapsed: None, pending_empty_input_quit: false, usage: None, + output_rate: OutputRate::default(), overlays: OverlayState::default(), preferences: TuiPreferences::default(), settings_defaults: TuiSettingsDefaults::default(), @@ -604,6 +607,10 @@ impl TuiState { self.usage = Some(usage); } + pub(crate) fn observe_output_rate(&mut self, event: &RuntimeEvent) { + self.output_rate.observe(event); + } + pub(crate) fn set_reasoning_effort_label(&mut self, label: Option) { self.reasoning_effort_label = label; } @@ -674,14 +681,22 @@ impl TuiState { } pub(crate) fn status_parts(&self) -> [String; 3] { - let usage = format_session_usage_full(self.usage.as_ref()); + let rate = self.output_rate.label(); + let usage = format_session_usage_full(self.usage.as_ref(), &rate); let model = self.model_status_label(); [self.workspace_root.display().to_string(), model, usage] } pub(crate) fn header_status_parts(&self, width: u16) -> [String; 3] { let model = self.model_status_label(); - format_header_status_parts(&self.workspace_root, &model, self.usage.as_ref(), width) + let rate = self.output_rate.label(); + format_header_status_parts( + &self.workspace_root, + &model, + self.usage.as_ref(), + &rate, + width, + ) } pub(crate) fn interaction_status_text(&self) -> String { diff --git a/crates/merry-cli/src/tui/status.rs b/crates/merry-cli/src/tui/status.rs index ac75c5ed..327ef17a 100644 --- a/crates/merry-cli/src/tui/status.rs +++ b/crates/merry-cli/src/tui/status.rs @@ -5,10 +5,10 @@ use unicode_width::{UnicodeWidthChar, UnicodeWidthStr}; const BRAND_AND_SEPARATORS_WIDTH: usize = 11; const MIN_WORKSPACE_WIDTH: usize = 8; -pub(crate) fn format_session_usage_full(usage: Option<&SessionUsage>) -> String { +pub(crate) fn format_session_usage_full(usage: Option<&SessionUsage>, rate: &str) -> String { usage - .map(format_session_usage) - .unwrap_or_else(SessionUsageDisplay::unavailable) + .map(|usage| format_session_usage(usage, rate)) + .unwrap_or_else(|| SessionUsageDisplay::unavailable(rate)) .full } @@ -16,12 +16,13 @@ pub(crate) fn format_header_status_parts( workspace: &Path, model: &str, usage: Option<&SessionUsage>, + rate: &str, width: u16, ) -> [String; 3] { let workspace = workspace.display().to_string(); let usage = usage - .map(format_session_usage) - .unwrap_or_else(SessionUsageDisplay::unavailable); + .map(|usage| format_session_usage(usage, rate)) + .unwrap_or_else(|| SessionUsageDisplay::unavailable(rate)); let width = usize::from(width); let minimum_workspace_width = display_width(&workspace).min(MIN_WORKSPACE_WIDTH); let model_width = display_width(model); @@ -34,7 +35,7 @@ pub(crate) fn format_header_status_parts( + display_width(candidate) <= width }) - .unwrap_or(usage.compact.as_str()) + .unwrap_or(usage.minimal.as_str()) .to_owned(); let remaining = width.saturating_sub(BRAND_AND_SEPARATORS_WIDTH + display_width(&usage)); @@ -53,14 +54,16 @@ struct SessionUsageDisplay { full: String, medium: String, compact: String, + minimal: String, } impl SessionUsageDisplay { - fn unavailable() -> Self { + fn unavailable(rate: &str) -> Self { Self { - full: "usage -".to_owned(), - medium: "usage -".to_owned(), - compact: "usage -".to_owned(), + full: format!("usage - · {rate}"), + medium: format!("usage - · {rate}"), + compact: format!("usage - · {rate}"), + minimal: "usage -".to_owned(), } } @@ -69,19 +72,22 @@ impl SessionUsageDisplay { self.full.as_str(), self.medium.as_str(), self.compact.as_str(), + self.minimal.as_str(), ] .into_iter() } } -fn format_session_usage(usage: &SessionUsage) -> SessionUsageDisplay { - let compact = format_context_pressure(usage); +fn format_session_usage(usage: &SessionUsage, rate: &str) -> SessionUsageDisplay { + let minimal = format_context_pressure(usage); + let compact = format!("{minimal} · {rate}"); let cache = format_cache_ratio(usage.last.input_tokens(), usage.last.cached_input_tokens()); - let medium = cache - .as_ref() - .map_or_else(|| compact.clone(), |cache| format!("{compact} · {cache}")); + let medium = cache.as_ref().map_or_else( + || compact.clone(), + |cache| format!("{minimal} · {cache} · {rate}"), + ); - let mut context_parts = vec![compact.clone()]; + let mut context_parts = vec![minimal.clone()]; if let Some(context) = usage.context { context_parts.push(format!( "win {} {}", @@ -93,7 +99,7 @@ fn format_session_usage(usage: &SessionUsage) -> SessionUsageDisplay { context_parts.push(cache); } let full = format!( - "{} | last in {} out {} | total {} tok", + "{} | last in {} out {} | total {} tok | {rate}", context_parts.join(" · "), format_token_count(usage.last.input_tokens()), format_token_count(usage.last.output_tokens()), @@ -104,35 +110,28 @@ fn format_session_usage(usage: &SessionUsage) -> SessionUsageDisplay { full, medium, compact, + minimal, } } fn format_context_pressure(usage: &SessionUsage) -> String { - let Some(compaction) = usage.compaction else { - return usage.context.map_or_else( - || "ctx -".to_owned(), - |context| { - format!( - "ctx in {}/{}", - format_token_count(usage.last.input_tokens()), - format_token_count(context.resolved_model_window_tokens) - ) - }, - ); - }; - - let current = compaction - .dynamic_body_estimated_tokens - .map(format_token_count) - .unwrap_or_else(|| "-".to_owned()); - if compaction.auto_compaction_enabled { - format!( - "ctx {current}/{}", - format_token_count(compaction.hard_water_tokens) - ) - } else { - format!("ctx {current} · compact off") + let mut context = usage.context.map_or_else( + || "ctx -".to_owned(), + |context| { + format!( + "ctx {}/{}", + format_token_count(usage.last.input_tokens()), + format_token_count(context.effective_window_tokens) + ) + }, + ); + if usage + .compaction + .is_some_and(|compaction| !compaction.auto_compaction_enabled) + { + context.push_str(" · compact off"); } + context } fn format_cache_ratio(input_tokens: u64, cached_input_tokens: Option) -> Option { diff --git a/crates/merry-cli/src/tui/tests.rs b/crates/merry-cli/src/tui/tests.rs index aa1ea098..59c23908 100644 --- a/crates/merry-cli/src/tui/tests.rs +++ b/crates/merry-cli/src/tui/tests.rs @@ -113,6 +113,7 @@ mod input_controls; mod layout; mod output_preview; +mod output_rate; mod patch_projection; diff --git a/crates/merry-cli/src/tui/tests/output_rate.rs b/crates/merry-cli/src/tui/tests/output_rate.rs new file mode 100644 index 00000000..db280765 --- /dev/null +++ b/crates/merry-cli/src/tui/tests/output_rate.rs @@ -0,0 +1,263 @@ +use super::{pending_call, source}; +use crate::tui::{ + keymap::Keymap, output_rate::OutputRate, projector::TuiProjector, render::render_to_text, + state::TuiState, theme::TuiTheme, +}; +use merry_core::{ + ContextWindowSource, ErrorInfo, InteractiveRunState, ModelOutputRate, ModelUsage, RuntimeEvent, + SessionUsage, UsageContextWindow, +}; +use std::time::Duration; + +fn rate_event(tokens: u64, seconds: u64, estimated: bool) -> RuntimeEvent { + RuntimeEvent::ModelOutputRateUpdated { + rate: Some(ModelOutputRate::new( + tokens, + Duration::from_secs(seconds), + if estimated { + merry_core::OutputTokenSource::Estimated + } else { + merry_core::OutputTokenSource::ProviderUsage + }, + )), + source: source(), + } +} + +#[test] +fn runtime_sample_drives_estimate_without_a_ui_stopwatch() { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(400, 8, true)); + assert_eq!(rate.label(), "≈50.0 tok/s"); + rate.observe(&rate_event(600, 10, true)); + assert_eq!(rate.label(), "≈60.0 tok/s"); +} + +#[test] +fn actual_usage_corrects_tokens_without_adding_delivery_or_completion_latency() { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(400, 8, true)); + rate.observe(&rate_event(2_400, 8, false)); + assert_eq!(rate.label(), "300.0 tok/s"); + rate.observe(&RuntimeEvent::AssistantMessageDelta { + delta: "late text".to_owned(), + source: source(), + }); + assert_eq!(rate.label(), "300.0 tok/s"); +} + +#[test] +fn unavailable_provider_timing_does_not_fall_back_to_text_or_running_time() { + let mut rate = OutputRate::default(); + rate.observe(&RuntimeEvent::StepStarted { source: source() }); + rate.observe(&RuntimeEvent::AssistantMessageDelta { + delta: "hello".to_owned(), + source: source(), + }); + assert_eq!(rate.label(), "- tok/s"); +} + +#[test] +fn restored_usage_does_not_create_a_rate_without_receive_timing() { + let mut rate = OutputRate::default(); + rate.observe(&RuntimeEvent::UsageUpdated { + usage: SessionUsage { + last: ModelUsage::new(100, 100), + total: ModelUsage::new(100, 100), + context: None, + compaction: None, + }, + source: source(), + }); + assert_eq!(rate.label(), "- tok/s"); +} + +#[test] +fn single_receive_instant_is_unavailable_but_zero_output_with_elapsed_time_is_valid() { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(100, 0, false)); + assert_eq!(rate.label(), "- tok/s"); + rate.observe(&rate_event(0, 2, false)); + assert_eq!(rate.label(), "0.0 tok/s"); +} + +#[test] +fn lifecycle_events_do_not_override_runtime_owned_samples() { + for event in [ + RuntimeEvent::StepStarted { source: source() }, + RuntimeEvent::SessionStarted { source: source() }, + RuntimeEvent::ModelRetryAttemptStarted { + attempt: 2, + max_attempts: 2, + source: source(), + }, + RuntimeEvent::CompactionStarted { source: source() }, + RuntimeEvent::InteractiveRunStateChanged { + state: InteractiveRunState::RunningModel, + }, + ] { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(100, 2, false)); + rate.observe(&event); + assert_eq!(rate.label(), "50.0 tok/s"); + } +} + +#[test] +fn tools_idle_cancellation_and_failure_do_not_extend_provider_receive_time() { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(100, 2, true)); + for event in [ + RuntimeEvent::ToolCallStarted { + call: pending_call("read", "read_text"), + source: source(), + }, + RuntimeEvent::InteractiveRunStateChanged { + state: InteractiveRunState::WaitingForInput, + }, + RuntimeEvent::RunCancelled { + diagnostic: ErrorInfo::new("cancelled", "cancelled").unwrap(), + source: source(), + }, + RuntimeEvent::RunFailed { + diagnostic: ErrorInfo::new("failed", "failed").unwrap(), + source: source(), + }, + RuntimeEvent::Closed, + ] { + rate.observe(&event); + assert_eq!(rate.label(), "≈50.0 tok/s"); + } +} + +#[test] +fn estimates_remain_available_when_a_provider_does_not_report_usage() { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(100, 2, true)); + rate.observe(&RuntimeEvent::StepCompleted { source: source() }); + assert_eq!(rate.label(), "≈50.0 tok/s"); +} + +#[test] +fn header_shows_tool_only_response_rate_and_preserves_context_on_narrow_screens() { + let mut state = TuiState::new( + "/repo".into(), + "gpt-test".to_owned(), + Keymap::default(), + TuiTheme::default(), + ); + let mut projector = TuiProjector::default(); + projector.apply(RuntimeEvent::StepStarted { source: source() }, &mut state); + projector.apply(rate_event(100, 2, false), &mut state); + projector.apply( + RuntimeEvent::UsageUpdated { + usage: SessionUsage { + total: ModelUsage::new(20_800, 100), + last: ModelUsage::new(20_800, 100), + context: Some(UsageContextWindow { + resolved_model_window_tokens: 272_000, + effective_window_tokens: 258_400, + source: ContextWindowSource::Fallback, + }), + compaction: None, + }, + source: source(), + }, + &mut state, + ); + for width in [160, 72] { + let rendered = render_to_text(&state, width, 16); + assert!(rendered.contains("50.0 tok/s")); + assert!(rendered.contains("ctx 20.8k/258.4k")); + } + let narrow = render_to_text(&state, 48, 16); + assert!(narrow.contains("ctx 20.8k/258.4k")); + assert!(!narrow.contains("50.0 tok/s")); + assert!( + state + .status_text() + .contains("last in 20.8k out 100 | total 20.9k tok") + ); +} + +#[test] +fn runtime_reset_does_not_hide_the_last_measurable_sample() { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(100, 2, false)); + rate.observe(&RuntimeEvent::ModelOutputRateUpdated { + rate: None, + source: source(), + }); + assert_eq!(rate.label(), "50.0 tok/s"); +} + +#[test] +fn backpressure_keeps_updating_the_rate_until_the_runtime_resets_it() { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(100, 2, false)); + assert_eq!(rate.label(), "50.0 tok/s"); + + for (tokens, token_source, label) in [ + ( + 400, + merry_core::OutputTokenSource::Estimated, + "≈100.0 tok/s", + ), + ( + 800, + merry_core::OutputTokenSource::ProviderUsage, + "≈200.0 tok/s", + ), + ] { + rate.observe(&RuntimeEvent::ModelOutputRateUpdated { + rate: Some( + ModelOutputRate::new(tokens, Duration::from_secs(4), token_source) + .with_timing_quality(merry_core::OutputTimingQuality::ConsumerLimited), + ), + source: source(), + }); + assert_eq!(rate.label(), label); + } + rate.observe(&RuntimeEvent::StepCompleted { source: source() }); + assert_eq!(rate.label(), "≈200.0 tok/s"); + + rate.observe(&RuntimeEvent::ModelOutputRateUpdated { + rate: None, + source: source(), + }); + assert_eq!(rate.label(), "≈200.0 tok/s"); + rate.observe(&rate_event(50, 0, true)); + assert_eq!(rate.label(), "≈200.0 tok/s"); + rate.observe(&rate_event(150, 1, false)); + assert_eq!(rate.label(), "150.0 tok/s"); +} + +#[test] +fn an_unmeasurable_new_sample_does_not_hide_the_last_measurable_sample() { + let mut rate = OutputRate::default(); + rate.observe(&rate_event(100, 2, false)); + rate.observe(&rate_event(2_400, 0, false)); + assert_eq!(rate.label(), "50.0 tok/s"); +} + +#[test] +fn timing_limitations_are_presented_as_estimates_not_missing_samples() { + let mut rate = OutputRate::default(); + for quality in [ + merry_core::OutputTimingQuality::PartialOutput, + merry_core::OutputTimingQuality::ConsumerLimited, + ] { + rate.observe(&RuntimeEvent::ModelOutputRateUpdated { + rate: Some( + ModelOutputRate::new( + 100, + Duration::from_secs(2), + merry_core::OutputTokenSource::ProviderUsage, + ) + .with_timing_quality(quality), + ), + source: source(), + }); + assert_eq!(rate.label(), "≈50.0 tok/s"); + } +} diff --git a/crates/merry-cli/src/tui/tests/status_usage.rs b/crates/merry-cli/src/tui/tests/status_usage.rs index f70d611f..17771455 100644 --- a/crates/merry-cli/src/tui/tests/status_usage.rs +++ b/crates/merry-cli/src/tui/tests/status_usage.rs @@ -57,7 +57,7 @@ fn projector_updates_queue_preview_and_usage_without_timeline_noise() { assert_eq!(state.queue_preview().next[0].text, "urgent"); assert!(state.timeline().is_empty()); - assert!(state.status_text().contains("ctx 20.2k/54.8k")); + assert!(state.status_text().contains("ctx 20k/60.8k")); assert!(state.status_text().contains("win 64k fallback")); assert!(state.status_text().contains("cache 90%")); assert!( @@ -67,6 +67,51 @@ fn projector_updates_queue_preview_and_usage_without_timeline_noise() { ); } +#[test] +fn context_shows_actual_input_over_safe_window_regardless_of_compaction_metadata() { + let compaction = CompactionUsageWindow { + auto_compaction_enabled: true, + dynamic_body_estimated_tokens: Some(1_700), + body_budget_tokens: 229_820, + soft_water_tokens: 218_940, + hard_water_tokens: 227_100, + }; + for compaction in [ + None, + Some(compaction), + Some(CompactionUsageWindow { + auto_compaction_enabled: false, + dynamic_body_estimated_tokens: None, + ..compaction + }), + ] { + let mut state = TuiState::new( + "/repo".into(), + "model".to_owned(), + Keymap::default(), + TuiTheme::default(), + ); + let last = ModelUsage::with_details(20_800, Some(20_592), 68, None, 20_868); + state.set_usage(SessionUsage { + total: last, + last, + context: Some(UsageContextWindow { + resolved_model_window_tokens: 272_000, + effective_window_tokens: 258_400, + source: ContextWindowSource::Fallback, + }), + compaction, + }); + let status = state.status_text(); + assert!(status.contains("ctx 20.8k/258.4k")); + assert!(status.contains("win 272k fallback · cache 99% | last in 20.8k out 68")); + assert_eq!( + status.contains("compact off"), + compaction.is_some_and(|budget| !budget.auto_compaction_enabled), + ); + } +} + #[test] fn narrow_header_preserves_context_pressure_before_secondary_usage() { let mut state = TuiState::new( @@ -94,7 +139,7 @@ fn narrow_header_preserves_context_pressure_before_secondary_usage() { let rendered = render_to_text(&state, 72, 16); - assert!(rendered.contains("ctx 20.2k/54.8k")); + assert!(rendered.contains("ctx 29k/60.8k")); assert!(!rendered.contains("total 2659.1k")); } @@ -125,7 +170,7 @@ fn narrow_header_counts_wide_characters_when_preserving_context() { let rendered = render_to_text(&state, 48, 16); - assert!(rendered.contains("ctx 20.2k/54.8k")); + assert!(rendered.contains("ctx 20k/60.8k")); } #[test] diff --git a/crates/merry-core/src/journal.rs b/crates/merry-core/src/journal.rs index e40ce18f..db37c19e 100644 --- a/crates/merry-core/src/journal.rs +++ b/crates/merry-core/src/journal.rs @@ -11,9 +11,11 @@ use serde::{Deserialize, Serialize}; /// Ordered runtime journal event. /// -/// Journal events are the runtime's durable execution log. They are suitable -/// for diagnostics, replay inspection, and runtime control, but they are not -/// the default SDK/UI event stream. +/// Ordered execution records and live observations. Transient payloads are +/// delivered to subscribers but excluded from replay and retained run results. +/// Replaceable telemetry, such as output rate, may be coalesced by a stream +/// consumer; semantic records retain their ordered delivery contract. This is +/// not the default SDK/UI event stream. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)] #[serde(deny_unknown_fields)] pub struct RuntimeJournalEvent { @@ -46,6 +48,12 @@ pub enum RuntimeJournalPayload { SessionStarted, /// A runtime step started. StepStarted, + /// Transient receive-side output throughput; not prompt or ledger state. + /// Consumers should treat this as replaceable telemetry rather than a + /// complete history of every intermediate sample. + ModelOutputRateUpdated { + rate: Option, + }, /// A model provider attempt started. ModelRetryAttemptStarted { /// 1-based attempt number. @@ -203,3 +211,14 @@ pub enum RuntimeJournalPayload { /// The runtime failed. Failed { diagnostic: ErrorInfo }, } + +impl RuntimeJournalPayload { + /// Whether this live-only observation is excluded from retained execution evidence. + #[must_use] + pub const fn is_transient(&self) -> bool { + matches!( + self, + Self::ModelOutputRateUpdated { .. } | Self::AssistantOutputDelta { .. } + ) + } +} diff --git a/crates/merry-core/src/lib.rs b/crates/merry-core/src/lib.rs index 1629288c..525d9002 100644 --- a/crates/merry-core/src/lib.rs +++ b/crates/merry-core/src/lib.rs @@ -6,6 +6,7 @@ pub mod event; pub mod evidence; pub mod id; pub mod journal; +pub mod output_rate; pub mod plan; pub mod runtime_event; pub mod schema; @@ -26,6 +27,7 @@ pub use id::{ ToolSourceId, TrajectoryRecordId, }; pub use journal::{RuntimeJournalEvent, RuntimeJournalPayload}; +pub use output_rate::{ModelOutputRate, OutputTimingQuality, OutputTokenSource}; pub use plan::{ CoordinatorDirectiveSnapshot, PlanActivationSource, PlanApprovalRequirementKind, PlanApprovalRequirementSnapshot, PlanApprovalRequirementStatus, PlanAttemptOutcome, diff --git a/crates/merry-core/src/output_rate.rs b/crates/merry-core/src/output_rate.rs new file mode 100644 index 00000000..2037f6c2 --- /dev/null +++ b/crates/merry-core/src/output_rate.rs @@ -0,0 +1,172 @@ +//! Provider-observed output throughput projected by the runtime. + +use schemars::JsonSchema; +use serde::{Deserialize, Serialize}; +use std::time::Duration; + +/// Source of the token count, independent of the quality of the receive interval. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "snake_case")] +pub enum OutputTokenSource { + /// UTF-8 byte estimate of observed output. + Estimated, + /// Total output tokens reported by the provider, including reported reasoning. + ProviderUsage, +} + +/// Limitations of the client-observed output window; never a server decode clock. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "snake_case")] +pub enum OutputTimingQuality { + /// First-to-last effective output received without consumer backpressure. + #[default] + ReceiveWindow, + /// Summarized, redacted, or otherwise unobserved reasoning makes the window partial. + PartialOutput, + /// Bounded delivery paused reading; the client receive rate includes consumer stalls. + ConsumerLimited, +} + +/// One output-rate observation with separate token and timing provenance. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +pub struct ModelOutputRate { + output_tokens: u64, + elapsed_nanos: u64, + token_source: OutputTokenSource, + timing_quality: OutputTimingQuality, +} + +impl ModelOutputRate { + /// Creates a sample using the provider's first-to-last output interval. + /// Zero-duration samples are valid but cannot yield a rate. + #[must_use] + pub fn new(output_tokens: u64, elapsed: Duration, token_source: OutputTokenSource) -> Self { + Self { + output_tokens, + elapsed_nanos: u64::try_from(elapsed.as_nanos()).unwrap_or(u64::MAX), + token_source, + timing_quality: OutputTimingQuality::ReceiveWindow, + } + } + + /// Attaches the receive-window limitation observed by the provider adapter. + #[must_use] + pub const fn with_timing_quality(mut self, quality: OutputTimingQuality) -> Self { + self.timing_quality = quality; + self + } + + /// Returns how the numerator was obtained. + #[must_use] + pub const fn token_source(self) -> OutputTokenSource { + self.token_source + } + + /// Returns whether the denominator covers only partial output or consumer stalls. + #[must_use] + pub const fn timing_quality(self) -> OutputTimingQuality { + self.timing_quality + } + + /// Whether token estimation or limited timing requires approximate presentation. + #[must_use] + pub const fn is_estimated(self) -> bool { + matches!(self.token_source, OutputTokenSource::Estimated) + || !matches!(self.timing_quality, OutputTimingQuality::ReceiveWindow) + } + + /// Output token count; includes reasoning when reported in the provider's output total. + #[must_use] + pub const fn output_tokens(self) -> u64 { + self.output_tokens + } + + /// Receive-side output duration, excluding time before first output and after last output. + #[must_use] + pub const fn elapsed(self) -> Duration { + Duration::from_nanos(self.elapsed_nanos) + } + + /// Observed client receive throughput, or None for zero duration. + /// Timing limitations retain an approximate rate; consult `is_estimated` and `timing_quality`. + #[must_use] + pub fn tokens_per_second(self) -> Option { + (!self.elapsed().is_zero()) + .then(|| self.output_tokens as f64 / self.elapsed().as_secs_f64()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn rate_uses_fractional_seconds_and_rejects_unmeasurable_intervals() { + assert_eq!( + ModelOutputRate::new( + 25, + Duration::from_millis(250), + crate::OutputTokenSource::ProviderUsage + ) + .tokens_per_second(), + Some(100.0) + ); + assert_eq!( + ModelOutputRate::new(25, Duration::ZERO, crate::OutputTokenSource::ProviderUsage) + .tokens_per_second(), + None + ); + assert_eq!( + ModelOutputRate::new( + 0, + Duration::from_secs(1), + crate::OutputTokenSource::ProviderUsage + ) + .tokens_per_second(), + Some(0.0) + ); + } + + #[test] + fn timing_limitations_downgrade_precision_without_discarding_receive_rate() { + for quality in [ + OutputTimingQuality::PartialOutput, + OutputTimingQuality::ConsumerLimited, + ] { + for source in [ + OutputTokenSource::Estimated, + OutputTokenSource::ProviderUsage, + ] { + let rate = ModelOutputRate::new(100, Duration::from_secs(2), source) + .with_timing_quality(quality); + assert_eq!(rate.tokens_per_second(), Some(50.0)); + assert!(rate.is_estimated()); + assert_eq!(rate.token_source(), source); + assert_eq!(rate.timing_quality(), quality); + + let untimed = + ModelOutputRate::new(100, Duration::ZERO, source).with_timing_quality(quality); + assert_eq!(untimed.tokens_per_second(), None); + } + } + } + + #[test] + fn rate_serialization_has_explicit_units_and_preserves_estimate_status() { + let rate = ModelOutputRate::new( + 50, + Duration::from_secs(2), + crate::OutputTokenSource::Estimated, + ); + let value = serde_json::to_value(rate).unwrap(); + assert_eq!( + value, + serde_json::json!({"output_tokens":50,"elapsed_nanos":2_000_000_000u64,"token_source":"estimated","timing_quality":"receive_window"}) + ); + assert_eq!( + serde_json::from_value::(value).unwrap(), + rate + ); + } +} diff --git a/crates/merry-core/src/runtime_event.rs b/crates/merry-core/src/runtime_event.rs index 66883d87..0e980eac 100644 --- a/crates/merry-core/src/runtime_event.rs +++ b/crates/merry-core/src/runtime_event.rs @@ -33,6 +33,11 @@ pub enum RuntimeEvent { usage: SessionUsage, source: RuntimeEventSource, }, + /// Live or usage-corrected throughput; None clears the active observation. + ModelOutputRateUpdated { + rate: Option, + source: RuntimeEventSource, + }, /// The assistant produced user-facing text. AssistantMessage { text: String, @@ -221,6 +226,17 @@ pub enum RuntimeEvent { Closed, } +impl RuntimeEvent { + /// Whether this live-only observation is excluded from retained run results. + #[must_use] + pub const fn is_transient(&self) -> bool { + matches!( + self, + Self::ModelOutputRateUpdated { .. } | Self::AssistantMessageDelta { .. } + ) + } +} + /// Pointer from a public event back to the journal position that produced it. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)] #[serde(deny_unknown_fields)] diff --git a/crates/merry-llm/src/event.rs b/crates/merry-llm/src/event.rs index 02af569f..5cab5d44 100644 --- a/crates/merry-llm/src/event.rs +++ b/crates/merry-llm/src/event.rs @@ -1,6 +1,6 @@ //! Normalized model event protocol. -use crate::{ModelResponse, ModelToolCall}; +use crate::{ModelOutputProgress, ModelResponse, ModelToolCall}; use schemars::JsonSchema; use serde::{Deserialize, Deserializer, Serialize}; @@ -10,6 +10,10 @@ use serde::{Deserialize, Deserializer, Serialize}; pub enum ModelEvent { /// Provider accepted the request and began producing output. Started, + /// Receive-side accounting without retaining text. None clears a failed attempt's sample. + OutputProgress { + progress: Option, + }, /// Incremental text output. OutputTextDelta { /// Delta text. @@ -31,9 +35,18 @@ pub enum ModelEvent { #[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)] enum ModelEventWire { Started {}, - OutputTextDelta { delta: String }, - ToolCallRequested { call: ModelToolCall }, - Completed { response: ModelResponse }, + OutputProgress { + progress: Option, + }, + OutputTextDelta { + delta: String, + }, + ToolCallRequested { + call: ModelToolCall, + }, + Completed { + response: ModelResponse, + }, } impl<'de> Deserialize<'de> for ModelEvent { @@ -44,6 +57,7 @@ impl<'de> Deserialize<'de> for ModelEvent { let wire = ModelEventWire::deserialize(deserializer)?; Ok(match wire { ModelEventWire::Started {} => Self::Started, + ModelEventWire::OutputProgress { progress } => Self::OutputProgress { progress }, ModelEventWire::OutputTextDelta { delta } => Self::OutputTextDelta { delta }, ModelEventWire::ToolCallRequested { call } => Self::ToolCallRequested { call }, ModelEventWire::Completed { response } => Self::Completed { response }, diff --git a/crates/merry-llm/src/lib.rs b/crates/merry-llm/src/lib.rs index db389a87..cf30a6ea 100644 --- a/crates/merry-llm/src/lib.rs +++ b/crates/merry-llm/src/lib.rs @@ -5,7 +5,9 @@ pub mod content; pub mod error; pub mod event; pub mod model_catalog; +pub mod output_progress; pub mod provider; +pub mod receive_stream; pub mod request; pub mod response; pub mod retry; @@ -22,7 +24,9 @@ pub use model_catalog::{ ModelCatalog, ModelCatalogEntry, ModelCatalogError, ModelCatalogErrorKind, ModelCatalogFuture, ModelCatalogProvider, }; +pub use output_progress::{ModelOutputProgress, OutputProgressTracker, StreamOutputKind}; pub use provider::{ModelEventStream, ModelProvider, ModelProviderFuture, ModelStreamContext}; +pub use receive_stream::receive_model_stream; pub use request::{ GenerationConfig, ModelInputItem, ModelMessage, ModelMessageRole, ModelName, ModelRequest, ModelResponseFormat, ModelStructuredOutputFormat, ParallelToolCalls, ReasoningEffort, diff --git a/crates/merry-llm/src/output_progress.rs b/crates/merry-llm/src/output_progress.rs new file mode 100644 index 00000000..4dbab7a1 --- /dev/null +++ b/crates/merry-llm/src/output_progress.rs @@ -0,0 +1,227 @@ +//! Bounded output accounting at the provider's stream-receive boundary. + +use merry_core::OutputTimingQuality; +use schemars::JsonSchema; +use serde::{Deserialize, Serialize}; +use std::time::{Duration, Instant}; + +/// Observed output bytes and the interval between first and last effective output. +/// Does not retain reasoning text, tool arguments, or a process-local clock in serialized data. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +pub struct ModelOutputProgress { + utf8_bytes: u64, + elapsed_nanos: u64, + reasoning_observed: bool, + timing_quality: OutputTimingQuality, +} + +impl ModelOutputProgress { + /// Creates a snapshot; zero elapsed time means throughput is not yet measurable. + #[must_use] + pub fn new(utf8_bytes: u64, elapsed: Duration) -> Self { + Self { + utf8_bytes, + elapsed_nanos: u64::try_from(elapsed.as_nanos()).unwrap_or(u64::MAX), + reasoning_observed: false, + timing_quality: OutputTimingQuality::ReceiveWindow, + } + } + + /// Total observed UTF-8 bytes across text, reasoning, and tool output. + #[must_use] + pub const fn utf8_bytes(self) -> u64 { + self.utf8_bytes + } + + /// Receive interval, excluding first-output latency and protocol completion tails. + #[must_use] + pub const fn elapsed(self) -> Duration { + Duration::from_nanos(self.elapsed_nanos) + } + + /// Marks whether raw reasoning, rather than just a summary, was observed. + #[must_use] + pub const fn with_reasoning_observed(mut self, observed: bool) -> Self { + self.reasoning_observed = observed; + self + } + + /// Records known receive-window limitations without retaining provider content. + #[must_use] + pub const fn with_timing_quality(mut self, quality: OutputTimingQuality) -> Self { + self.timing_quality = quality; + self + } + + /// Whether raw reasoning contributed to this observation. + #[must_use] + pub const fn reasoning_observed(self) -> bool { + self.reasoning_observed + } + + /// Quality of the client receive interval. + #[must_use] + pub const fn timing_quality(self) -> OutputTimingQuality { + self.timing_quality + } +} + +/// Classification of model-generated stream content; metadata is not output. +#[derive(Debug, Clone, Copy)] +pub enum StreamOutputKind { + /// Visible assistant text or tool name/argument fragments. + Content, + /// Provider-exposed reasoning text. + Reasoning, + /// A reasoning summary, used only when raw reasoning is unavailable. + ReasoningSummary, +} + +/// Per-attempt receive-side accounting shared by protocol adapters. +/// A new instance is required for every attempt. Empty deltas do not move either endpoint. +#[derive(Debug, Default)] +pub struct OutputProgressTracker { + first: Option, + last: Option, + content_bytes: u64, + reasoning_bytes: u64, + summary_bytes: u64, + incomplete_reasoning: bool, +} + +impl OutputProgressTracker { + /// Marks opaque or otherwise unobservable reasoning without counting it as visible text. + pub fn mark_incomplete_reasoning(&mut self) { + self.incomplete_reasoning = true; + } + + /// Accounts for a validated output fragment at the time its network data was received. + /// Raw reasoning takes precedence over summaries to avoid counting both representations. + pub fn observe(&mut self, kind: StreamOutputKind, text: &str, received_at: Instant) { + if text.is_empty() { + return; + } + let bytes = u64::try_from(text.len()).unwrap_or(u64::MAX); + match kind { + StreamOutputKind::Content => { + self.content_bytes = self.content_bytes.saturating_add(bytes) + } + StreamOutputKind::Reasoning => { + self.reasoning_bytes = self.reasoning_bytes.saturating_add(bytes) + } + StreamOutputKind::ReasoningSummary => { + self.summary_bytes = self.summary_bytes.saturating_add(bytes); + if self.reasoning_bytes > 0 { + return; + } + } + } + self.first.get_or_insert(received_at); + self.last = Some(received_at); + } + + /// Latest observation; None until effective output has arrived. + #[must_use] + pub fn snapshot(&self) -> Option { + let elapsed = self.last?.saturating_duration_since(self.first?); + let reasoning = if self.reasoning_bytes > 0 { + self.reasoning_bytes + } else { + self.summary_bytes + }; + let quality = + if self.incomplete_reasoning || (self.reasoning_bytes == 0 && self.summary_bytes > 0) { + OutputTimingQuality::PartialOutput + } else { + OutputTimingQuality::ReceiveWindow + }; + Some( + ModelOutputProgress::new(self.content_bytes.saturating_add(reasoning), elapsed) + .with_reasoning_observed(self.reasoning_bytes > 0) + .with_timing_quality(quality), + ) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn output_progress_counts_utf8_independent_of_fragment_boundaries() { + let now = Instant::now(); + let mut tracker = OutputProgressTracker::default(); + tracker.observe(StreamOutputKind::Content, "", now); + assert_eq!(tracker.snapshot(), None); + tracker.observe( + StreamOutputKind::Reasoning, + "思", + now + Duration::from_secs(10), + ); + tracker.observe( + StreamOutputKind::Reasoning, + "考", + now + Duration::from_secs(10), + ); + tracker.observe( + StreamOutputKind::Content, + "ok", + now + Duration::from_secs(12), + ); + tracker.observe( + StreamOutputKind::Content, + "", + now + Duration::from_secs(100), + ); + assert_eq!( + tracker.snapshot(), + Some(ModelOutputProgress::new(8, Duration::from_secs(2)).with_reasoning_observed(true)) + ); + assert_eq!(OutputProgressTracker::default().snapshot(), None); + } + + #[test] + fn output_progress_prefers_raw_reasoning_and_ignores_late_summary_timing() { + let now = Instant::now(); + let mut tracker = OutputProgressTracker::default(); + tracker.observe(StreamOutputKind::ReasoningSummary, "summary", now); + tracker.observe( + StreamOutputKind::Reasoning, + "actual reasoning", + now + Duration::from_secs(1), + ); + tracker.observe( + StreamOutputKind::Content, + "body", + now + Duration::from_secs(2), + ); + tracker.observe( + StreamOutputKind::ReasoningSummary, + "duplicate", + now + Duration::from_secs(30), + ); + assert_eq!( + tracker.snapshot(), + Some( + ModelOutputProgress::new(20, Duration::from_secs(2)).with_reasoning_observed(true) + ) + ); + } + + #[test] + fn output_progress_round_trips_without_serializing_clock_or_text() { + let event = crate::ModelEvent::OutputProgress { + progress: Some(ModelOutputProgress::new(7, Duration::from_millis(250))), + }; + let json = serde_json::to_value(&event).unwrap(); + assert_eq!( + json, + serde_json::json!({"type":"output_progress","progress":{"utf8_bytes":7,"elapsed_nanos":250_000_000,"reasoning_observed":false,"timing_quality":"receive_window"}}) + ); + assert_eq!( + serde_json::from_value::(json).unwrap(), + event + ); + } +} diff --git a/crates/merry-llm/src/receive_stream.rs b/crates/merry-llm/src/receive_stream.rs new file mode 100644 index 00000000..dc109900 --- /dev/null +++ b/crates/merry-llm/src/receive_stream.rs @@ -0,0 +1,292 @@ +//! Owned receive task with bounded semantic delivery and replaceable telemetry. + +use crate::{ModelError, ModelEvent, ModelEventStream, ModelOutputProgress, ProviderErrorKind}; +use futures_core::Stream; +use futures_util::{StreamExt, stream}; +use merry_core::OutputTimingQuality; +use std::{ + future::Future, + pin::Pin, + task::{Context, Poll}, + time::Duration, +}; +use tokio::{ + sync::{mpsc, watch}, + task::JoinHandle, +}; +use tokio_util::sync::CancellationToken; + +const EVENT_CAPACITY: usize = 32; +const PROGRESS_INTERVAL: Duration = Duration::from_millis(100); + +/// Starts an owned receive task for one provider attempt. +/// Semantic events remain ordered and bounded; live progress is latest-only and +/// published at most every 100 ms of observed output, plus resets and completion. +/// Slow semantic delivery marks timing as consumer-limited rather than reporting +/// a misleading rate. Completion awaits the task; dropping the stream cancels +/// and aborts it, dropping its input stream and network resources. +#[must_use] +pub fn receive_model_stream( + source: ModelEventStream, + token: &CancellationToken, +) -> ModelEventStream { + let token = token.child_token(); + let (sender, receiver) = mpsc::channel(EVENT_CAPACITY); + let (progress_sender, progress_receiver) = watch::channel(ProgressUpdate::default()); + let worker_token = token.clone(); + let worker = tokio::spawn(receive(source, sender, progress_sender, worker_token)); + let progress = stream::unfold(progress_receiver, |mut receiver| async move { + receiver.changed().await.ok()?; + let update = *receiver.borrow_and_update(); + Some((update, receiver)) + }); + Box::pin(ReceiveStream { + receiver, + progress: Box::pin(progress.fuse()), + worker: Some(worker), + token, + last_revision: 0, + prefer_progress: false, + terminal: None, + closed: false, + done: false, + }) +} + +#[derive(Clone, Copy, Default)] +struct ProgressUpdate { + revision: u64, + progress: Option, +} + +enum Delivery { + Event(Result), + FinalProgress(ProgressUpdate), +} + +#[derive(Default)] +struct ProgressPublisher { + latest: ProgressUpdate, + last_published: Option, + consumer_limited: bool, +} + +impl ProgressPublisher { + fn observe( + &mut self, + progress: Option, + sender: &watch::Sender, + ) { + if progress.is_none() { + self.consumer_limited = false; + } + self.latest.revision = self.latest.revision.saturating_add(1); + self.latest.progress = progress.map(|progress| { + if self.consumer_limited { + progress.with_timing_quality(OutputTimingQuality::ConsumerLimited) + } else { + progress + } + }); + let elapsed = progress.map(ModelOutputProgress::elapsed); + if elapsed.is_none_or(|elapsed| { + self.last_published + .is_none_or(|last| elapsed.saturating_sub(last) >= PROGRESS_INTERVAL) + }) { + self.last_published = elapsed; + sender.send_replace(self.latest); + } + } + + fn mark_consumer_limited(&mut self, sender: &watch::Sender) { + if !self.consumer_limited { + self.consumer_limited = true; + if let Some(progress) = self.latest.progress { + self.latest.revision = self.latest.revision.saturating_add(1); + self.latest.progress = + Some(progress.with_timing_quality(OutputTimingQuality::ConsumerLimited)); + sender.send_replace(self.latest); + } + } + } +} + +async fn receive( + mut source: ModelEventStream, + sender: mpsc::Sender, + progress_sender: watch::Sender, + token: CancellationToken, +) -> Result<(), ModelError> { + let mut progress = ProgressPublisher::default(); + loop { + tokio::task::consume_budget().await; + let item = tokio::select! { + biased; + () = token.cancelled() => return Err(ModelError::Cancelled), + () = sender.closed() => return Ok(()), + item = source.next() => item, + }; + if let Some(Ok(ModelEvent::OutputProgress { progress: sample })) = item { + progress.observe(sample, &progress_sender); + continue; + } + let terminal = matches!( + item, + None | Some(Err(_)) | Some(Ok(ModelEvent::Completed { .. })) + ); + if terminal && progress.latest.revision != 0 { + progress_sender.send_replace(progress.latest); + if !send(&sender, &token, Delivery::FinalProgress(progress.latest)).await? { + return Ok(()); + } + } + let Some(item) = item else { + return Ok(()); + }; + let delivery = Delivery::Event(item); + if terminal { + let _ = send(&sender, &token, delivery).await?; + return Ok(()); + } + match sender.try_send(delivery) { + Ok(()) => {} + Err(mpsc::error::TrySendError::Closed(_)) => return Ok(()), + Err(mpsc::error::TrySendError::Full(delivery)) => { + progress.mark_consumer_limited(&progress_sender); + if !send(&sender, &token, delivery).await? { + return Ok(()); + } + } + } + } +} + +async fn send( + sender: &mpsc::Sender, + token: &CancellationToken, + event: Delivery, +) -> Result { + tokio::select! { + biased; + () = token.cancelled() => Err(ModelError::Cancelled), + result = sender.send(event) => Ok(result.is_ok()), + } +} + +struct ReceiveStream { + receiver: mpsc::Receiver, + progress: Pin + Send>>, + worker: Option>>, + token: CancellationToken, + last_revision: u64, + prefer_progress: bool, + terminal: Option>, + closed: bool, + done: bool, +} + +impl ReceiveStream { + fn observation(&mut self, update: ProgressUpdate) -> Option> { + if update.revision <= self.last_revision { + return None; + } + self.last_revision = update.revision; + self.prefer_progress = false; + Some(Ok(ModelEvent::OutputProgress { + progress: update.progress, + })) + } + + fn poll_progress( + &mut self, + context: &mut Context<'_>, + ) -> Poll>> { + while let Poll::Ready(update) = self.progress.as_mut().poll_next(context) { + let Some(update) = update else { + return Poll::Ready(None); + }; + if let Some(event) = self.observation(update) { + return Poll::Ready(Some(event)); + } + } + Poll::Pending + } + + fn poll_finish( + &mut self, + context: &mut Context<'_>, + ) -> Poll>> { + let Some(worker) = self.worker.as_mut() else { + return Poll::Ready(None); + }; + let result = std::task::ready!(Pin::new(worker).poll(context)); + self.worker = None; + self.done = true; + if self.token.is_cancelled() { + return Poll::Ready(Some(Err(ModelError::Cancelled))); + } + Poll::Ready(match result { + Ok(Ok(())) => self.terminal.take(), + Ok(Err(error)) => Some(Err(error)), + Err(_) => Some(Err(ModelError::provider( + ProviderErrorKind::Other, + "model stream receive task failed", + ))), + }) + } +} + +impl Stream for ReceiveStream { + type Item = Result; + + fn poll_next(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll> { + let state = self.get_mut(); + loop { + if state.done { + return Poll::Ready(None); + } + if state.token.is_cancelled() || state.closed || state.terminal.is_some() { + return state.poll_finish(context); + } + if state.prefer_progress + && let Poll::Ready(Some(event)) = state.poll_progress(context) + { + return Poll::Ready(Some(event)); + } + match state.receiver.poll_recv(context) { + Poll::Ready(Some(Delivery::Event(event))) => { + if matches!(event, Err(_) | Ok(ModelEvent::Completed { .. })) { + state.terminal = Some(event); + continue; + } + state.prefer_progress = true; + return Poll::Ready(Some(event)); + } + Poll::Ready(Some(Delivery::FinalProgress(update))) => { + if let Some(event) = state.observation(update) { + return Poll::Ready(Some(event)); + } + } + Poll::Ready(None) => state.closed = true, + Poll::Pending => { + if let Poll::Ready(Some(event)) = state.poll_progress(context) { + return Poll::Ready(Some(event)); + } + return Poll::Pending; + } + } + } + } +} + +impl Drop for ReceiveStream { + fn drop(&mut self) { + self.token.cancel(); + if let Some(worker) = self.worker.take() { + worker.abort(); + } + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/merry-llm/src/receive_stream/tests.rs b/crates/merry-llm/src/receive_stream/tests.rs new file mode 100644 index 00000000..79d55952 --- /dev/null +++ b/crates/merry-llm/src/receive_stream/tests.rs @@ -0,0 +1,226 @@ +use super::*; +use crate::{FinishReason, ModelOutput, ModelResponse, Usage}; +use tokio::sync::oneshot; + +fn progress(index: u64) -> Result { + Ok(ModelEvent::OutputProgress { + progress: Some(ModelOutputProgress::new( + index * 4, + Duration::from_millis(index), + )), + }) +} + +fn completed() -> Result { + Ok(ModelEvent::Completed { + response: ModelResponse::new( + vec![ModelOutput::text("done")], + FinishReason::Stop, + Some(Usage::new(1, 100_000)), + ), + }) +} + +async fn signal(receiver: oneshot::Receiver<()>) { + tokio::time::timeout(Duration::from_secs(2), receiver) + .await + .unwrap() + .unwrap(); +} + +#[tokio::test] +async fn paused_consumer_does_not_retain_or_block_a_long_reasoning_stream() { + let (finished, received) = oneshot::channel(); + let source = stream::iter((1..=100_000).map(progress)).chain(stream::once(async move { + let _ = finished.send(()); + completed() + })); + let output = receive_model_stream(Box::pin(source), &CancellationToken::new()); + signal(received).await; + let events: Vec<_> = output.collect().await; + let samples: Vec<_> = events + .iter() + .filter_map(|event| match event { + Ok(ModelEvent::OutputProgress { + progress: Some(progress), + }) => Some(*progress), + _ => None, + }) + .collect(); + assert!( + samples.len() <= 2, + "intermediate observations must be replaced, not queued" + ); + assert_eq!( + samples.last(), + Some(&ModelOutputProgress::new(400_000, Duration::from_secs(100))) + ); + assert!(matches!( + events.last(), + Some(Ok(ModelEvent::Completed { .. })) + )); +} + +#[tokio::test] +async fn fast_consumer_still_receives_rate_limited_progress_and_the_final_sample() { + let source = stream::iter((1..=2_000).map(progress).chain([completed()])); + let events: Vec<_> = receive_model_stream(Box::pin(source), &CancellationToken::new()) + .collect() + .await; + let samples: Vec<_> = events + .iter() + .filter_map(|event| match event { + Ok(ModelEvent::OutputProgress { + progress: Some(progress), + }) => Some(*progress), + _ => None, + }) + .collect(); + assert!(samples.len() <= 22); + assert_eq!( + samples.last(), + Some(&ModelOutputProgress::new(8_000, Duration::from_secs(2))) + ); +} + +#[tokio::test] +async fn semantic_backpressure_is_bounded_and_marks_timing_without_losing_text() { + let (full, received) = oneshot::channel(); + let mut full = Some(full); + let source = stream::iter([progress(1)]) + .chain(stream::iter((0..EVENT_CAPACITY * 2).map(move |index| { + if index == EVENT_CAPACITY { + let _ = full.take().unwrap().send(()); + } + Ok(ModelEvent::OutputTextDelta { + delta: "x".to_owned(), + }) + }))) + .chain(stream::iter([progress(2_000), completed()])); + let output = receive_model_stream(Box::pin(source), &CancellationToken::new()); + signal(received).await; + let events: Vec<_> = output.collect().await; + let text: String = events + .iter() + .filter_map(|event| match event { + Ok(ModelEvent::OutputTextDelta { delta }) => Some(delta.as_str()), + _ => None, + }) + .collect(); + assert_eq!(text, "x".repeat(EVENT_CAPACITY * 2)); + let sample = events + .iter() + .rev() + .find_map(|event| match event { + Ok(ModelEvent::OutputProgress { + progress: Some(progress), + }) => Some(*progress), + _ => None, + }) + .unwrap(); + assert_eq!( + sample.timing_quality(), + OutputTimingQuality::ConsumerLimited + ); + assert_eq!(sample.utf8_bytes(), 8_000); +} + +struct PendingSource(Option>); + +impl Stream for PendingSource { + type Item = Result; + + fn poll_next(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll> { + Poll::Pending + } +} + +impl Drop for PendingSource { + fn drop(&mut self) { + if let Some(sender) = self.0.take() { + let _ = sender.send(()); + } + } +} + +#[tokio::test] +async fn dropping_the_stream_stops_its_source_without_cancelling_the_parent() { + let token = CancellationToken::new(); + let (dropped, received) = oneshot::channel(); + let output = receive_model_stream(Box::pin(PendingSource(Some(dropped))), &token); + drop(output); + signal(received).await; + assert!(!token.is_cancelled()); +} + +#[tokio::test] +async fn cancellation_awaits_source_cleanup_before_returning_the_terminal_error() { + let token = CancellationToken::new(); + let (dropped, received) = oneshot::channel(); + let mut output = receive_model_stream(Box::pin(PendingSource(Some(dropped))), &token); + token.cancel(); + assert!(matches!( + output.next().await, + Some(Err(ModelError::Cancelled)) + )); + signal(received).await; + assert!(output.next().await.is_none()); +} + +#[tokio::test] +async fn exceptional_worker_shutdown_is_reported_instead_of_silent_eof() { + let source = stream::once(async { panic!("fixture receive failure") }); + let mut output = receive_model_stream(Box::pin(source), &CancellationToken::new()); + let error = output.next().await.unwrap().unwrap_err(); + assert_eq!(error.kind(), ProviderErrorKind::Other); + assert!(output.next().await.is_none()); +} + +#[tokio::test] +async fn independent_attempts_do_not_reuse_progress_or_timing_quality() { + for _ in 0..2 { + let source = stream::iter([progress(1_000), completed()]); + let events: Vec<_> = receive_model_stream(Box::pin(source), &CancellationToken::new()) + .collect() + .await; + assert!(events.iter().any(|event| matches!(event, Ok(ModelEvent::OutputProgress { progress: Some(progress) }) if progress.timing_quality() == OutputTimingQuality::ReceiveWindow))); + } +} + +#[tokio::test] +async fn receive_clock_is_sampled_before_a_delayed_consumer_resumes() { + use std::sync::{ + Arc, + atomic::{AtomicU64, Ordering}, + }; + let clock = Arc::new(AtomicU64::new(0)); + let receive_clock = Arc::clone(&clock); + let (release, gate) = oneshot::channel(); + let (read, read_done) = oneshot::channel(); + let source = stream::iter([progress(0)]) + .chain(stream::once(async move { + gate.await.unwrap(); + let observed = progress(receive_clock.load(Ordering::SeqCst)); + let _ = read.send(()); + observed + })) + .chain(stream::iter([completed()])); + let output = receive_model_stream(Box::pin(source), &CancellationToken::new()); + clock.store(2_000, Ordering::SeqCst); + release.send(()).unwrap(); + signal(read_done).await; + clock.store(100_000, Ordering::SeqCst); + let events: Vec<_> = output.collect().await; + let last = events + .iter() + .rev() + .find_map(|event| match event { + Ok(ModelEvent::OutputProgress { + progress: Some(progress), + }) => Some(progress), + _ => None, + }) + .unwrap(); + assert_eq!(last.elapsed(), Duration::from_secs(2)); + assert_eq!(last.timing_quality(), OutputTimingQuality::ReceiveWindow); +} diff --git a/crates/merry-llm/src/retry.rs b/crates/merry-llm/src/retry.rs index 14db9162..883d5077 100644 --- a/crates/merry-llm/src/retry.rs +++ b/crates/merry-llm/src/retry.rs @@ -358,6 +358,7 @@ async fn run_retry_stream( result = inner.stream_model(request.clone(), context.stream_context()) => result, }; + let mut observed_progress = false; let error = match setup { Ok(mut stream) => { let mut committed = false; @@ -375,6 +376,7 @@ async fn run_retry_stream( match item { Some(Ok(ModelEvent::Started)) => {} Some(Ok(event)) => { + observed_progress |= matches!(event, ModelEvent::OutputProgress { .. }); committed |= commits_output(&event); let completed = matches!(event, ModelEvent::Completed { .. }); if !send_stream_item(&sender, &token, Ok(event)).await { @@ -411,6 +413,18 @@ async fn run_retry_stream( return; }; + if observed_progress + && !send_stream_item( + &sender, + &token, + Ok(ModelEvent::OutputProgress { progress: None }), + ) + .await + { + send_cancelled_if_open(&sender); + return; + } + tokio::select! { biased; () = token.cancelled() => { diff --git a/crates/merry-llm/src/retry/tests.rs b/crates/merry-llm/src/retry/tests.rs index f98592ad..a25b6755 100644 --- a/crates/merry-llm/src/retry/tests.rs +++ b/crates/merry-llm/src/retry/tests.rs @@ -171,6 +171,56 @@ async fn retrying_provider_does_not_retry_after_visible_output() { )); } +#[tokio::test] +async fn telemetry_preserves_retry_policy_and_clears_failed_attempt_progress_in_stream_order() { + let progress = ModelEvent::OutputProgress { + progress: Some(crate::ModelOutputProgress::new(4, Duration::ZERO)), + }; + let inner = AttemptScriptProvider::new(vec![ + vec![ + Ok(progress.clone()), + Err(ModelError::provider( + ProviderErrorKind::Unavailable, + "stream broke", + )), + ], + vec![Ok(completed("recovered"))], + ]); + let retrying = RetryingModelProvider::new( + Arc::new(inner.clone()), + ModelRetryPolicy::new( + true, + 2, + Duration::from_millis(1), + Duration::from_millis(1), + Duration::from_secs(10), + false, + ) + .unwrap(), + ); + let events = retrying + .stream_model(request(), ModelStreamContext::default()) + .await + .unwrap() + .collect::>() + .await; + assert!( + events + .iter() + .any(|event| event.as_ref().is_ok_and(|event| event == &progress)) + ); + assert_eq!( + events.into_iter().collect::, _>>().unwrap(), + vec![ + ModelEvent::Started, + progress, + ModelEvent::OutputProgress { progress: None }, + completed("recovered"), + ] + ); + assert_eq!(inner.attempts_remaining(), 0); +} + #[tokio::test] async fn retrying_provider_retries_before_visible_output_with_one_started_event() { let inner = AttemptScriptProvider::new(vec![ diff --git a/crates/merry-provider-anthropic/src/parse.rs b/crates/merry-provider-anthropic/src/parse.rs index c4bf4c55..25909b46 100644 --- a/crates/merry-provider-anthropic/src/parse.rs +++ b/crates/merry-provider-anthropic/src/parse.rs @@ -7,10 +7,11 @@ use crate::{ }; use merry_core::ToolName; use merry_llm::{ - FinishReason, ModelEvent, ModelOutput, ModelResponse, ModelToolCall, ModelToolCallId, - ToolArguments, Usage, + FinishReason, ModelEvent, ModelOutput, ModelOutputProgress, ModelResponse, ModelToolCall, + ModelToolCallId, OutputProgressTracker, StreamOutputKind, ToolArguments, Usage, }; use std::collections::BTreeMap; +use std::time::Instant; pub(crate) struct AnthropicStreamParser { aggregate_text: String, @@ -20,6 +21,7 @@ pub(crate) struct AnthropicStreamParser { finish_reason: Option, usage: UsageAccumulator, completed: bool, + output: OutputProgressTracker, } impl AnthropicStreamParser { @@ -32,12 +34,26 @@ impl AnthropicStreamParser { finish_reason: None, usage: UsageAccumulator::default(), completed: false, + output: OutputProgressTracker::default(), } } + #[cfg(test)] pub(crate) fn parse_sse_line( &mut self, raw_line: &str, + ) -> Result, AnthropicProviderError> { + self.parse_sse_line_at(raw_line, Instant::now()) + } + + pub(crate) fn output_progress(&self) -> Option { + self.output.snapshot() + } + + pub(crate) fn parse_sse_line_at( + &mut self, + raw_line: &str, + received_at: Instant, ) -> Result, AnthropicProviderError> { let line = raw_line.trim_end_matches(['\r', '\n']); if line.is_empty() || line.starts_with(':') { @@ -66,7 +82,14 @@ impl AnthropicStreamParser { "failed to parse Anthropic stream event: {error}" )) })?; - self.parse_event(event) + self.parse_event(event, received_at) + } + + /// Thinking text can be a summary rather than a complete token-level trace. + fn observe_thinking(&mut self, thinking: &str, received_at: Instant) { + self.output.mark_incomplete_reasoning(); + self.output + .observe(StreamOutputKind::Reasoning, thinking, received_at); } pub(crate) fn finish(&self) -> Result<(), AnthropicProviderError> { @@ -82,6 +105,7 @@ impl AnthropicStreamParser { fn parse_event( &mut self, event: AnthropicStreamEvent, + received_at: Instant, ) -> Result, AnthropicProviderError> { match event { AnthropicStreamEvent::MessageStart { message } => { @@ -94,7 +118,17 @@ impl AnthropicStreamParser { index, content_block, } => match content_block { + AnthropicContentBlockStart::RedactedThinking => { + self.output.mark_incomplete_reasoning(); + Ok(Vec::new()) + } + AnthropicContentBlockStart::Thinking { thinking } => { + self.observe_thinking(&thinking, received_at); + Ok(Vec::new()) + } AnthropicContentBlockStart::Text { text } if !text.is_empty() => { + self.output + .observe(StreamOutputKind::Content, &text, received_at); self.aggregate_text.push_str(&text); Ok(vec![ModelEvent::OutputTextDelta { delta: text }]) } @@ -112,6 +146,10 @@ impl AnthropicStreamParser { )) })? }; + self.output + .observe(StreamOutputKind::Content, &name, received_at); + self.output + .observe(StreamOutputKind::Content, &initial_input, received_at); if self .tool_buffers .insert( @@ -132,7 +170,13 @@ impl AnthropicStreamParser { } }, AnthropicStreamEvent::ContentBlockDelta { index, delta } => match delta { + AnthropicContentBlockDelta::ThinkingDelta { thinking } => { + self.observe_thinking(&thinking, received_at); + Ok(Vec::new()) + } AnthropicContentBlockDelta::TextDelta { text } if !text.is_empty() => { + self.output + .observe(StreamOutputKind::Content, &text, received_at); self.aggregate_text.push_str(&text); Ok(vec![ModelEvent::OutputTextDelta { delta: text }]) } @@ -148,6 +192,8 @@ impl AnthropicStreamParser { })? .input .push_str(&partial_json); + self.output + .observe(StreamOutputKind::Content, &partial_json, received_at); Ok(Vec::new()) } }, diff --git a/crates/merry-provider-anthropic/src/provider.rs b/crates/merry-provider-anthropic/src/provider.rs index 515cf734..f64c2e27 100644 --- a/crates/merry-provider-anthropic/src/provider.rs +++ b/crates/merry-provider-anthropic/src/provider.rs @@ -9,11 +9,17 @@ use merry_llm::{ ModelProviderFuture, ModelRequest, ModelStreamContext, ProviderErrorKind, }; use serde_json::Value; -use std::{collections::VecDeque, time::Duration}; +use std::{ + collections::VecDeque, + time::{Duration, Instant}, +}; use tracing::Instrument; const USER_AGENT_VALUE: &str = concat!("merry/", env!("CARGO_PKG_VERSION")); +#[cfg(test)] +mod output_progress_tests; + /// Config-backed Anthropic Messages provider. #[derive(Debug, Clone)] pub struct AnthropicProvider { @@ -111,7 +117,10 @@ impl ModelProvider for AnthropicProvider { AnthropicEventStreamState::new(response, token, stream_span), |state| async move { state.next_item().await }, ); - Ok(Box::pin(event_stream) as ModelEventStream) + Ok(merry_llm::receive_model_stream( + Box::pin(event_stream), + context.cancellation_token(), + )) } .instrument(span), ) @@ -175,7 +184,7 @@ impl AnthropicEventStreamState { }; match chunk { Ok(Some(chunk)) => { - if let Err(error) = self.events.parse_bytes(&chunk) { + if let Err(error) = self.events.parse_bytes_at(&chunk, Instant::now()) { self.done = true; return Some(( Err(add_stream_endpoint_context( @@ -219,6 +228,7 @@ struct AnthropicEventStreamEvents { parser: AnthropicStreamParser, line_buffer: Vec, pending: VecDeque, + last_received_at: Option, } impl AnthropicEventStreamEvents { @@ -227,6 +237,7 @@ impl AnthropicEventStreamEvents { parser: AnthropicStreamParser::new(), line_buffer: Vec::new(), pending: VecDeque::from([ModelEvent::Started]), + last_received_at: None, } } @@ -234,21 +245,35 @@ impl AnthropicEventStreamEvents { self.pending.pop_front() } - fn parse_bytes(&mut self, bytes: &[u8]) -> Result<(), AnthropicProviderError> { + fn parse_bytes_at( + &mut self, + bytes: &[u8], + received_at: Instant, + ) -> Result<(), AnthropicProviderError> { + self.last_received_at = Some(received_at); for byte in bytes { self.line_buffer.push(*byte); if *byte == b'\n' { - self.parse_buffered_line()?; + self.parse_buffered_line(received_at)?; } } Ok(()) } - fn parse_buffered_line(&mut self) -> Result<(), AnthropicProviderError> { + fn parse_buffered_line(&mut self, received_at: Instant) -> Result<(), AnthropicProviderError> { let line = std::str::from_utf8(&self.line_buffer).map_err(|error| { AnthropicProviderError::protocol(format!("stream line is not UTF-8: {error}")) })?; - self.pending.extend(self.parser.parse_sse_line(line)?); + let previous = self.parser.output_progress(); + let events = self.parser.parse_sse_line_at(line, received_at)?; + if let Some(progress) = self.parser.output_progress() + && Some(progress) != previous + { + self.pending.push_back(ModelEvent::OutputProgress { + progress: Some(progress), + }); + } + self.pending.extend(events); self.line_buffer.clear(); Ok(()) } @@ -257,7 +282,7 @@ impl AnthropicEventStreamEvents { &mut self, ) -> Result, AnthropicProviderError> { if !self.line_buffer.is_empty() { - self.parse_buffered_line()?; + self.parse_buffered_line(self.last_received_at.unwrap_or_else(Instant::now))?; } self.parser.finish()?; Ok(self.pop_pending()) diff --git a/crates/merry-provider-anthropic/src/provider/output_progress_tests.rs b/crates/merry-provider-anthropic/src/provider/output_progress_tests.rs new file mode 100644 index 00000000..c68e343c --- /dev/null +++ b/crates/merry-provider-anthropic/src/provider/output_progress_tests.rs @@ -0,0 +1,142 @@ +use super::AnthropicEventStreamEvents; +use merry_llm::{ModelEvent, ModelOutputProgress}; +use std::time::{Duration, Instant}; + +#[test] +fn thinking_text_and_tool_json_share_receive_timing_without_signature_or_usage_tails() { + let now = Instant::now(); + let mut events = AnthropicEventStreamEvents::new(); + for (seconds, data) in [ + ( + 1, + r#"{"type":"message_start","message":{"usage":{"input_tokens":100,"output_tokens":0}}}"#, + ), + ( + 2, + r#"{"type":"content_block_start","index":0,"content_block":{"type":"thinking","thinking":""}}"#, + ), + ( + 10, + r#"{"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"思考"}}"#, + ), + ( + 11, + r#"{"type":"content_block_delta","index":0,"delta":{"type":"signature_delta","signature":"not model output"}}"#, + ), + (12, r#"{"type":"content_block_stop","index":0}"#), + ( + 13, + r#"{"type":"content_block_start","index":1,"content_block":{"type":"text","text":"ok"}}"#, + ), + ( + 14, + r#"{"type":"content_block_delta","index":1,"delta":{"type":"text_delta","text":"ay"}}"#, + ), + (15, r#"{"type":"content_block_stop","index":1}"#), + ( + 16, + r#"{"type":"content_block_start","index":2,"content_block":{"type":"tool_use","id":"call-1","name":"read","input":{}}}"#, + ), + ( + 18, + r#"{"type":"content_block_delta","index":2,"delta":{"type":"input_json_delta","partial_json":"{}"}}"#, + ), + (19, r#"{"type":"content_block_stop","index":2}"#), + ( + 40, + r#"{"type":"message_delta","delta":{"stop_reason":"tool_use"},"usage":{"output_tokens":2400}}"#, + ), + (50, r#"{"type":"message_stop"}"#), + ] { + events + .parse_bytes_at( + format!("data: {data}\n").as_bytes(), + now + Duration::from_secs(seconds), + ) + .unwrap(); + } + let mut samples = Vec::new(); + let mut visible = String::new(); + let mut usage = None; + while let Some(event) = events.pop_pending() { + match event { + ModelEvent::OutputProgress { + progress: Some(progress), + } => samples.push(progress), + ModelEvent::OutputTextDelta { delta } => visible.push_str(&delta), + ModelEvent::Completed { response } => usage = response.usage(), + _ => {} + } + } + assert_eq!(visible, "okay"); + assert_eq!(samples.len(), 5); + assert_eq!( + samples.first(), + Some( + &ModelOutputProgress::new(6, Duration::ZERO) + .with_reasoning_observed(true) + .with_timing_quality(merry_core::OutputTimingQuality::PartialOutput) + ) + ); + assert_eq!( + samples.last(), + Some( + &ModelOutputProgress::new(16, Duration::from_secs(8)) + .with_reasoning_observed(true) + .with_timing_quality(merry_core::OutputTimingQuality::PartialOutput) + ) + ); + assert_eq!(usage.unwrap().output_tokens, 2400); +} + +#[test] +fn initial_thinking_counts_but_redacted_data_and_empty_deltas_do_not() { + let now = Instant::now(); + let mut events = AnthropicEventStreamEvents::new(); + for (seconds, data) in [ + ( + 0, + r#"{"type":"content_block_start","index":0,"content_block":{"type":"redacted_thinking","data":"opaque"}}"#, + ), + (1, r#"{"type":"ping"}"#), + ( + 2, + r#"{"type":"content_block_start","index":1,"content_block":{"type":"thinking","thinking":"plan"}}"#, + ), + ( + 3, + r#"{"type":"content_block_delta","index":1,"delta":{"type":"thinking_delta","thinking":""}}"#, + ), + ( + 10, + r#"{"type":"content_block_start","index":2,"content_block":{"type":"text","text":"done"}}"#, + ), + ] { + events + .parse_bytes_at( + format!("data: {data}\n").as_bytes(), + now + Duration::from_secs(seconds), + ) + .unwrap(); + } + let mut samples = Vec::new(); + while let Some(event) = events.pop_pending() { + if let ModelEvent::OutputProgress { + progress: Some(progress), + } = event + { + samples.push(progress); + } + } + assert_eq!( + samples, + vec![ + ModelOutputProgress::new(4, Duration::ZERO) + .with_reasoning_observed(true) + .with_timing_quality(merry_core::OutputTimingQuality::PartialOutput), + ModelOutputProgress::new(8, Duration::from_secs(8)) + .with_reasoning_observed(true) + .with_timing_quality(merry_core::OutputTimingQuality::PartialOutput) + ] + ); +} diff --git a/crates/merry-provider-anthropic/src/wire.rs b/crates/merry-provider-anthropic/src/wire.rs index e2caf377..2398a6ce 100644 --- a/crates/merry-provider-anthropic/src/wire.rs +++ b/crates/merry-provider-anthropic/src/wire.rs @@ -120,6 +120,10 @@ pub(crate) struct AnthropicMessageStart { #[derive(Debug, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] pub(crate) enum AnthropicContentBlockStart { + RedactedThinking, + Thinking { + thinking: String, + }, Text { text: String, }, @@ -135,6 +139,9 @@ pub(crate) enum AnthropicContentBlockStart { #[derive(Debug, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] pub(crate) enum AnthropicContentBlockDelta { + ThinkingDelta { + thinking: String, + }, TextDelta { text: String, }, diff --git a/crates/merry-provider-openai/src/chat_completions/parse.rs b/crates/merry-provider-openai/src/chat_completions/parse.rs index bfdb025a..d7240612 100644 --- a/crates/merry-provider-openai/src/chat_completions/parse.rs +++ b/crates/merry-provider-openai/src/chat_completions/parse.rs @@ -2,10 +2,11 @@ use super::wire::{ChatChunk, ChatUsage}; use crate::OpenAiProviderError; use merry_core::ToolName; use merry_llm::{ - FinishReason, ModelEvent, ModelOutput, ModelResponse, ModelToolCall, ModelToolCallId, - ToolArguments, Usage, + FinishReason, ModelEvent, ModelOutput, ModelOutputProgress, ModelResponse, ModelToolCall, + ModelToolCallId, OutputProgressTracker, StreamOutputKind, ToolArguments, Usage, }; use std::collections::BTreeMap; +use std::time::Instant; pub(crate) struct ChatStreamParser { aggregate_text: String, @@ -14,6 +15,7 @@ pub(crate) struct ChatStreamParser { finish_reason: Option, usage: Option, completed: bool, + output: OutputProgressTracker, } impl ChatStreamParser { @@ -25,12 +27,26 @@ impl ChatStreamParser { finish_reason: None, usage: None, completed: false, + output: OutputProgressTracker::default(), } } + #[cfg(test)] pub(crate) fn parse_sse_line( &mut self, raw_line: &str, + ) -> Result, OpenAiProviderError> { + self.parse_sse_line_at(raw_line, Instant::now()) + } + + pub(crate) fn output_progress(&self) -> Option { + self.output.snapshot() + } + + pub(crate) fn parse_sse_line_at( + &mut self, + raw_line: &str, + received_at: Instant, ) -> Result, OpenAiProviderError> { let line = raw_line.trim_end_matches(['\r', '\n']); if line.is_empty() || line.starts_with(':') { @@ -65,7 +81,7 @@ impl ChatStreamParser { "failed to parse Chat Completions stream chunk: {error}" )) })?; - self.parse_chunk(chunk) + self.parse_chunk(chunk, received_at) } pub(crate) fn finish(&self) -> Result<(), OpenAiProviderError> { @@ -78,7 +94,11 @@ impl ChatStreamParser { } } - fn parse_chunk(&mut self, chunk: ChatChunk) -> Result, OpenAiProviderError> { + fn parse_chunk( + &mut self, + chunk: ChatChunk, + received_at: Instant, + ) -> Result, OpenAiProviderError> { if chunk.choices.len() > 1 { return Err(OpenAiProviderError::protocol( "Chat Completions stream returned multiple choices", @@ -98,13 +118,34 @@ impl ChatStreamParser { } let mut events = Vec::new(); + if let Some(reasoning) = choice + .delta + .reasoning_content + .as_deref() + .filter(|text| !text.is_empty()) + .or(choice.delta.reasoning.as_deref()) + { + self.output + .observe(StreamOutputKind::Reasoning, reasoning, received_at); + } if let Some(content) = choice.delta.content && !content.is_empty() { + self.output + .observe(StreamOutputKind::Content, &content, received_at); self.aggregate_text.push_str(&content); events.push(ModelEvent::OutputTextDelta { delta: content }); } for delta in choice.delta.tool_calls { + if let Some(function) = delta.function.as_ref() { + for fragment in [function.name.as_deref(), function.arguments.as_deref()] + .into_iter() + .flatten() + { + self.output + .observe(StreamOutputKind::Content, fragment, received_at); + } + } self.tool_buffers .entry(delta.index) .or_default() diff --git a/crates/merry-provider-openai/src/chat_completions/wire.rs b/crates/merry-provider-openai/src/chat_completions/wire.rs index 64407924..531da9eb 100644 --- a/crates/merry-provider-openai/src/chat_completions/wire.rs +++ b/crates/merry-provider-openai/src/chat_completions/wire.rs @@ -123,6 +123,8 @@ pub(crate) struct ChatChoice { #[derive(Debug, Default, Deserialize)] pub(crate) struct ChatDelta { pub(crate) content: Option, + pub(crate) reasoning_content: Option, + pub(crate) reasoning: Option, #[serde(default)] pub(crate) tool_calls: Vec, } diff --git a/crates/merry-provider-openai/src/parse.rs b/crates/merry-provider-openai/src/parse.rs index 20cd8bae..0f322ffb 100644 --- a/crates/merry-provider-openai/src/parse.rs +++ b/crates/merry-provider-openai/src/parse.rs @@ -10,12 +10,14 @@ use crate::{ }; use merry_core::ToolName; use merry_llm::{ - FinishDetail, FinishReason, ModelEvent, ModelOutput, ModelResponse, ModelToolCall, - ModelToolCallId, ProviderErrorKind, ToolArguments, Usage, + FinishDetail, FinishReason, ModelEvent, ModelOutput, ModelOutputProgress, ModelResponse, + ModelToolCall, ModelToolCallId, OutputProgressTracker, ProviderErrorKind, StreamOutputKind, + ToolArguments, Usage, }; use serde::Deserialize; use serde_json::Value; use std::collections::BTreeMap; +use std::time::Instant; const OPENAI_COMPATIBLE_HEARTBEAT_EVENT: &str = "ping"; @@ -61,6 +63,7 @@ pub(crate) struct ResponsesStreamParser { tool_call_buffers: BTreeMap, tool_calls: Vec, completed: bool, + output: OutputProgressTracker, } impl ResponsesStreamParser { @@ -70,12 +73,25 @@ impl ResponsesStreamParser { tool_call_buffers: BTreeMap::new(), tool_calls: Vec::new(), completed: false, + output: OutputProgressTracker::default(), } } pub(crate) fn parse_sse_line( &mut self, raw_line: &str, + ) -> Result, OpenAiProviderError> { + self.parse_sse_line_at(raw_line, Instant::now()) + } + + pub(crate) fn output_progress(&self) -> Option { + self.output.snapshot() + } + + pub(crate) fn parse_sse_line_at( + &mut self, + raw_line: &str, + received_at: Instant, ) -> Result, OpenAiProviderError> { let line = raw_line.trim_end_matches(['\r', '\n']); if line.is_empty() || line.starts_with(':') { @@ -142,7 +158,7 @@ impl ResponsesStreamParser { )) })?; - let result = self.parse_event(event); + let result = self.parse_event(event, received_at); tracing::debug!( event_type = metadata.event_type.unwrap_or(""), sequence_number = ?metadata.sequence_number, @@ -169,25 +185,40 @@ impl ResponsesStreamParser { fn parse_event( &mut self, event: ResponsesStreamEvent, + received_at: Instant, ) -> Result, OpenAiProviderError> { match event { ResponsesStreamEvent::Created | ResponsesStreamEvent::Other => Ok(Vec::new()), + ResponsesStreamEvent::ReasoningTextDelta { delta } => { + self.output + .observe(StreamOutputKind::Reasoning, &delta, received_at); + Ok(Vec::new()) + } + ResponsesStreamEvent::ReasoningSummaryTextDelta { delta } => { + self.output + .observe(StreamOutputKind::ReasoningSummary, &delta, received_at); + Ok(Vec::new()) + } ResponsesStreamEvent::OutputTextDelta { delta } => { if delta.is_empty() { return Ok(Vec::new()); } + self.output + .observe(StreamOutputKind::Content, &delta, received_at); self.aggregate_text.push_str(&delta); Ok(vec![ModelEvent::OutputTextDelta { delta }]) } ResponsesStreamEvent::OutputItemAdded { output_index, item } => { - self.merge_output_item(output_index, item)?; + self.merge_output_item(output_index, item, received_at)?; Ok(Vec::new()) } ResponsesStreamEvent::FunctionCallArgumentsDelta { output_index, delta, } => { + self.output + .observe(StreamOutputKind::Content, &delta, received_at); self.tool_call_buffers .entry(output_index) .or_default() @@ -199,16 +230,19 @@ impl ResponsesStreamParser { output_index, arguments, } => { - self.tool_call_buffers - .entry(output_index) - .or_default() - .set_arguments(arguments)?; + let buffer = self.tool_call_buffers.entry(output_index).or_default(); + let unseen = buffer.arguments.is_empty(); + buffer.set_arguments(arguments)?; + if unseen { + self.output + .observe(StreamOutputKind::Content, &buffer.arguments, received_at); + } Ok(Vec::new()) } ResponsesStreamEvent::OutputItemDone { output_index, item } => match item { ResponsesStreamOutputItem::Other => Ok(Vec::new()), item @ ResponsesStreamOutputItem::FunctionCall { .. } => { - self.merge_output_item(output_index, item)?; + self.merge_output_item(output_index, item, received_at)?; let buffer = self .tool_call_buffers .remove(&output_index) @@ -264,17 +298,28 @@ impl ResponsesStreamParser { &mut self, output_index: u64, item: ResponsesStreamOutputItem, + received_at: Instant, ) -> Result<(), OpenAiProviderError> { match item { ResponsesStreamOutputItem::FunctionCall { call_id, name, arguments, - } => self - .tool_call_buffers - .entry(output_index) - .or_default() - .merge(call_id, name, arguments), + } => { + let buffer = self.tool_call_buffers.entry(output_index).or_default(); + let unseen_name = buffer.name.is_none(); + let unseen_arguments = buffer.arguments.is_empty(); + buffer.merge(call_id, name, arguments)?; + if unseen_name && let Some(name) = buffer.name.as_deref() { + self.output + .observe(StreamOutputKind::Content, name, received_at); + } + if unseen_arguments { + self.output + .observe(StreamOutputKind::Content, &buffer.arguments, received_at); + } + Ok(()) + } ResponsesStreamOutputItem::Other => Ok(()), } } diff --git a/crates/merry-provider-openai/src/provider.rs b/crates/merry-provider-openai/src/provider.rs index 99ed5403..f3ae6d12 100644 --- a/crates/merry-provider-openai/src/provider.rs +++ b/crates/merry-provider-openai/src/provider.rs @@ -12,7 +12,7 @@ use merry_llm::{ }; use serde_json::Value; use std::collections::VecDeque; -use std::time::Duration; +use std::time::{Duration, Instant}; use tracing::Instrument; mod request; @@ -139,7 +139,10 @@ impl ModelProvider for OpenAiProvider { ); let event_stream: ModelEventStream = Box::pin(event_stream); tracing::debug!("openai event stream created"); - Ok(event_stream) + Ok(merry_llm::receive_model_stream( + event_stream, + context.cancellation_token(), + )) } .instrument(stream_span), ) @@ -210,11 +213,12 @@ impl OpenAiEventStreamState { match chunk { Ok(Some(chunk)) => { + let received_at = Instant::now(); tracing::trace!( chunk_byte_length = chunk.len(), "openai stream chunk received" ); - if let Err(error) = self.events.parse_bytes(chunk.as_ref()) { + if let Err(error) = self.events.parse_bytes_at(chunk.as_ref(), received_at) { tracing::debug!("openai stream protocol error"); self.done = true; return Some(( @@ -281,6 +285,7 @@ struct OpenAiEventStreamEvents { parser: OpenAiStreamParser, line_buffer: Vec, pending: VecDeque, + last_received_at: Option, } impl OpenAiEventStreamEvents { @@ -296,6 +301,7 @@ impl OpenAiEventStreamEvents { }, line_buffer: Vec::new(), pending: VecDeque::from([ModelEvent::Started]), + last_received_at: None, } } @@ -303,11 +309,21 @@ impl OpenAiEventStreamEvents { self.pending.pop_front() } + #[cfg(test)] fn parse_bytes(&mut self, bytes: &[u8]) -> Result<(), OpenAiProviderError> { + self.parse_bytes_at(bytes, Instant::now()) + } + + fn parse_bytes_at( + &mut self, + bytes: &[u8], + received_at: Instant, + ) -> Result<(), OpenAiProviderError> { + self.last_received_at = Some(received_at); for (byte_offset, byte) in bytes.iter().enumerate() { self.line_buffer.push(*byte); if *byte == b'\n' - && let Err(error) = self.parse_buffered_line() + && let Err(error) = self.parse_buffered_line(received_at) { tracing::debug!( chunk_byte_length = bytes.len(), @@ -322,18 +338,27 @@ impl OpenAiEventStreamEvents { Ok(()) } - fn parse_buffered_line(&mut self) -> Result<(), OpenAiProviderError> { + fn parse_buffered_line(&mut self, received_at: Instant) -> Result<(), OpenAiProviderError> { let line = std::str::from_utf8(&self.line_buffer).map_err(|error| { OpenAiProviderError::protocol(format!("stream line is not valid UTF-8: {error}")) })?; - self.pending.extend(self.parser.parse_sse_line(line)?); + let previous = self.parser.output_progress(); + let events = self.parser.parse_sse_line_at(line, received_at)?; + if let Some(progress) = self.parser.output_progress() + && Some(progress) != previous + { + self.pending.push_back(ModelEvent::OutputProgress { + progress: Some(progress), + }); + } + self.pending.extend(events); self.line_buffer.clear(); Ok(()) } fn finish_stream(&mut self) -> Result<(), OpenAiProviderError> { if !self.line_buffer.is_empty() { - self.parse_buffered_line()?; + self.parse_buffered_line(self.last_received_at.unwrap_or_else(Instant::now))?; } self.parser.finish() @@ -351,10 +376,21 @@ enum OpenAiStreamParser { } impl OpenAiStreamParser { - fn parse_sse_line(&mut self, line: &str) -> Result, OpenAiProviderError> { + fn parse_sse_line_at( + &mut self, + line: &str, + received_at: Instant, + ) -> Result, OpenAiProviderError> { + match self { + Self::Responses(parser) => parser.parse_sse_line_at(line, received_at), + Self::ChatCompletions(parser) => parser.parse_sse_line_at(line, received_at), + } + } + + fn output_progress(&self) -> Option { match self { - Self::Responses(parser) => parser.parse_sse_line(line), - Self::ChatCompletions(parser) => parser.parse_sse_line(line), + Self::Responses(parser) => parser.output_progress(), + Self::ChatCompletions(parser) => parser.output_progress(), } } @@ -447,6 +483,7 @@ fn model_event_category(event: &ModelEvent) -> &'static str { match event { ModelEvent::Started => "started", ModelEvent::OutputTextDelta { .. } => "output_text_delta", + ModelEvent::OutputProgress { .. } => "output_progress", ModelEvent::ToolCallRequested { .. } => "tool_call_requested", ModelEvent::Completed { .. } => "completed", } diff --git a/crates/merry-provider-openai/src/provider/tests.rs b/crates/merry-provider-openai/src/provider/tests.rs index 1d5772e1..d9ff8dd9 100644 --- a/crates/merry-provider-openai/src/provider/tests.rs +++ b/crates/merry-provider-openai/src/provider/tests.rs @@ -1,4 +1,5 @@ mod errors; +mod output_progress; mod request; mod retry; mod stream; diff --git a/crates/merry-provider-openai/src/provider/tests/output_progress.rs b/crates/merry-provider-openai/src/provider/tests/output_progress.rs new file mode 100644 index 00000000..bb28697f --- /dev/null +++ b/crates/merry-provider-openai/src/provider/tests/output_progress.rs @@ -0,0 +1,217 @@ +use crate::{OpenAiProtocol, provider::OpenAiEventStreamEvents}; +use merry_llm::{ModelEvent, ModelOutputProgress}; +use std::time::{Duration, Instant}; + +fn receive(events: &mut OpenAiEventStreamEvents, now: Instant, seconds: u64, data: &str) { + events + .parse_bytes_at( + format!("data: {data}\n").as_bytes(), + now + Duration::from_secs(seconds), + ) + .unwrap(); +} + +fn drain(events: &mut OpenAiEventStreamEvents) -> (Vec, String) { + let mut samples = Vec::new(); + let mut visible = String::new(); + while let Some(event) = events.pop_pending() { + match event { + ModelEvent::OutputProgress { + progress: Some(progress), + } => samples.push(progress), + ModelEvent::OutputTextDelta { delta } => visible.push_str(&delta), + _ => {} + } + } + (samples, visible) +} + +#[test] +fn chat_counts_ds_reasoning_text_and_tool_fragments_without_start_or_done_latency() { + let now = Instant::now(); + let mut events = OpenAiEventStreamEvents::new(OpenAiProtocol::ChatCompletions); + for (seconds, data) in [ + ( + 1, + r#"{"choices":[{"index":0,"delta":{"role":"assistant","content":null}}]}"#, + ), + ( + 10, + r#"{"choices":[{"index":0,"delta":{"reasoning_content":"思考"}}]}"#, + ), + ( + 12, + r#"{"choices":[{"index":0,"delta":{"reasoning_content":"abcd","reasoning":"duplicate"}}]}"#, + ), + (14, r#"{"choices":[{"index":0,"delta":{"content":"ok"}}]}"#), + ( + 16, + r#"{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call-1","function":{"name":"read","arguments":"{\"path\":"}}]}}]}"#, + ), + ( + 18, + r#"{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"a\"}"}}]}}]}"#, + ), + ( + 20, + r#"{"choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}"#, + ), + ( + 22, + r#"{"choices":[],"usage":{"prompt_tokens":20000,"completion_tokens":2400,"completion_tokens_details":{"reasoning_tokens":2000}}}"#, + ), + (30, "[DONE]"), + ] { + receive(&mut events, now, seconds, data); + } + events.finish_stream().unwrap(); + let (samples, visible) = drain(&mut events); + assert_eq!( + samples.first(), + Some(&ModelOutputProgress::new(6, Duration::ZERO).with_reasoning_observed(true)) + ); + assert_eq!( + samples.last(), + Some(&ModelOutputProgress::new(28, Duration::from_secs(8)).with_reasoning_observed(true)) + ); + assert_eq!(samples.len(), 5); + assert_eq!(visible, "ok"); +} + +#[test] +fn responses_reasoning_and_tool_output_ignore_summary_duplicates_and_done_snapshots() { + let now = Instant::now(); + let mut events = OpenAiEventStreamEvents::new(OpenAiProtocol::Responses); + for (seconds, data) in [ + (1, r#"{"type":"response.created"}"#), + ( + 10, + r#"{"type":"response.reasoning_text.delta","delta":"思考"}"#, + ), + ( + 11, + r#"{"type":"response.reasoning_summary_text.delta","delta":"summary"}"#, + ), + ( + 12, + r#"{"type":"response.output_item.added","output_index":1,"item":{"type":"function_call","call_id":"call-1","name":"read","arguments":""}}"#, + ), + ( + 14, + r#"{"type":"response.function_call_arguments.delta","output_index":1,"delta":"{}"}"#, + ), + ( + 50, + r#"{"type":"response.function_call_arguments.done","output_index":1,"arguments":"{}"}"#, + ), + ( + 60, + r#"{"type":"response.output_item.done","output_index":1,"item":{"type":"function_call","call_id":"call-1","name":"read","arguments":"{}"}}"#, + ), + ( + 70, + r#"{"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":100,"output_tokens":70}}}"#, + ), + ] { + receive(&mut events, now, seconds, data); + } + events.finish_stream().unwrap(); + let (samples, visible) = drain(&mut events); + assert_eq!(samples.len(), 3); + assert_eq!( + samples.last(), + Some(&ModelOutputProgress::new(12, Duration::from_secs(4)).with_reasoning_observed(true)) + ); + assert!(visible.is_empty()); +} + +#[test] +fn responses_summary_is_an_estimation_proxy_when_raw_reasoning_is_not_exposed() { + let now = Instant::now(); + let mut events = OpenAiEventStreamEvents::new(OpenAiProtocol::Responses); + receive( + &mut events, + now, + 10, + r#"{"type":"response.reasoning_summary_text.delta","delta":"plan"}"#, + ); + receive( + &mut events, + now, + 12, + r#"{"type":"response.output_text.delta","delta":"done"}"#, + ); + receive( + &mut events, + now, + 30, + r#"{"type":"response.output_text.done","text":"done"}"#, + ); + let (samples, visible) = drain(&mut events); + assert_eq!( + samples.last(), + Some( + &ModelOutputProgress::new(8, Duration::from_secs(2)) + .with_timing_quality(merry_core::OutputTimingQuality::PartialOutput) + ) + ); + assert_eq!(visible, "done"); +} + +#[test] +fn transport_coalescing_and_late_consumer_do_not_invent_an_output_interval() { + let now = Instant::now(); + let mut events = OpenAiEventStreamEvents::new(OpenAiProtocol::ChatCompletions); + events.parse_bytes_at(b"data: {\"choices\":[{\"index\":0,\"delta\":{\"reasoning\":\"plan\"}}]}\ndata: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"done\"}}]}\n", now).unwrap(); + let (samples, visible) = drain(&mut events); + assert_eq!( + samples.last(), + Some(&ModelOutputProgress::new(8, Duration::ZERO).with_reasoning_observed(true)) + ); + assert_eq!(visible, "done"); +} + +#[test] +fn empty_deltas_heartbeats_and_usage_do_not_start_a_measurement() { + for protocol in [OpenAiProtocol::Responses, OpenAiProtocol::ChatCompletions] { + let mut events = OpenAiEventStreamEvents::new(protocol); + events + .parse_bytes_at(b": heartbeat\n", Instant::now()) + .unwrap(); + let data = match protocol { + OpenAiProtocol::Responses => r#"{"type":"response.output_text.delta","delta":""}"#, + OpenAiProtocol::ChatCompletions => { + r#"{"choices":[{"index":0,"delta":{"reasoning_content":"","content":""}}]}"# + } + }; + receive(&mut events, Instant::now(), 60, data); + assert!(drain(&mut events).0.is_empty()); + } +} + +#[test] +fn split_utf8_lines_are_timed_on_receipt_and_eof_preserves_the_last_chunk_time() { + let now = Instant::now(); + let mut events = OpenAiEventStreamEvents::new(OpenAiProtocol::Responses); + let line = "data: {\"type\":\"response.reasoning_text.delta\",\"delta\":\"思考\"}\n"; + let split = line.find('思').unwrap() + 1; + events + .parse_bytes_at(&line.as_bytes()[..split], now) + .unwrap(); + events + .parse_bytes_at(&line.as_bytes()[split..], now + Duration::from_secs(10)) + .unwrap(); + events + .parse_bytes_at( + b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}", + now + Duration::from_secs(12), + ) + .unwrap(); + assert!(events.finish_stream().is_err()); + let (samples, visible) = drain(&mut events); + assert_eq!( + samples.last(), + Some(&ModelOutputProgress::new(8, Duration::from_secs(2)).with_reasoning_observed(true)) + ); + assert_eq!(visible, "ok"); +} diff --git a/crates/merry-provider-openai/src/provider/tests/stream.rs b/crates/merry-provider-openai/src/provider/tests/stream.rs index a96bb2d2..bc96d243 100644 --- a/crates/merry-provider-openai/src/provider/tests/stream.rs +++ b/crates/merry-provider-openai/src/provider/tests/stream.rs @@ -175,6 +175,15 @@ fn stream_state_emits_completed_from_final_usage_line_without_trailing_newline() events .parse_bytes(b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"Done\"}\n") .expect("text delta should parse"); + assert_eq!( + events.pop_pending(), + Some(ModelEvent::OutputProgress { + progress: Some(merry_llm::ModelOutputProgress::new( + 4, + std::time::Duration::ZERO + )), + }) + ); assert_eq!( events.pop_pending(), Some(ModelEvent::OutputTextDelta { @@ -210,6 +219,15 @@ fn eof_finalization_returns_completed_from_final_usage_line_without_trailing_new events .parse_bytes(b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"Done\"}\n") .expect("text delta should parse"); + assert_eq!( + events.pop_pending(), + Some(ModelEvent::OutputProgress { + progress: Some(merry_llm::ModelOutputProgress::new( + 4, + std::time::Duration::ZERO + )), + }) + ); assert_eq!( events.pop_pending(), Some(ModelEvent::OutputTextDelta { diff --git a/crates/merry-provider-openai/src/wire.rs b/crates/merry-provider-openai/src/wire.rs index caacd2d3..cbdaf528 100644 --- a/crates/merry-provider-openai/src/wire.rs +++ b/crates/merry-provider-openai/src/wire.rs @@ -256,6 +256,10 @@ pub(crate) enum ResponsesStreamEvent { Created, #[serde(rename = "response.output_text.delta")] OutputTextDelta { delta: String }, + #[serde(rename = "response.reasoning_text.delta")] + ReasoningTextDelta { delta: String }, + #[serde(rename = "response.reasoning_summary_text.delta")] + ReasoningSummaryTextDelta { delta: String }, #[serde(rename = "response.output_item.added")] OutputItemAdded { output_index: u64, diff --git a/crates/merry-runtime/src/agent_loop/helpers.rs b/crates/merry-runtime/src/agent_loop/helpers.rs index 8e2ef1ca..0f77040d 100644 --- a/crates/merry-runtime/src/agent_loop/helpers.rs +++ b/crates/merry-runtime/src/agent_loop/helpers.rs @@ -43,7 +43,10 @@ pub(crate) fn trace_loop_error(session_id: &str, model_turns_run: usize, source: pub(crate) async fn collect_step_events( stream: RuntimeJournalEventStream, ) -> Vec { - stream.collect().await + stream + .filter(|event| std::future::ready(!event.payload.is_transient())) + .collect() + .await } pub(crate) async fn final_assistant_output_from_step( @@ -128,9 +131,12 @@ pub(crate) async fn publish_journal_event( events: &mut Vec, event: RuntimeJournalEvent, ) -> Result<(), RuntimeError> { - let projected = projector.project(event.clone(), runtime).await?; + let retained = (!event.payload.is_transient()).then(|| event.clone()); + let projected = projector.project(event, runtime).await?; - events.push(event); + if let Some(event) = retained { + events.push(event); + } if let Some(projected) = projected && sender diff --git a/crates/merry-runtime/src/agent_loop/producer.rs b/crates/merry-runtime/src/agent_loop/producer.rs index 4b3cdfd9..6b2ce394 100644 --- a/crates/merry-runtime/src/agent_loop/producer.rs +++ b/crates/merry-runtime/src/agent_loop/producer.rs @@ -142,7 +142,9 @@ pub(crate) async fn run_agent_loop_stream_producer( let mut step_events = Vec::new(); tokio::pin!(stream); while let Some(event) = stream.next().await { - step_events.push(event.clone()); + if !event.payload.is_transient() { + step_events.push(event.clone()); + } publish_journal_event(&runtime, &mut projector, &sender, &mut events, event) .await .map_err(|source| { diff --git a/crates/merry-runtime/src/agent_loop/types.rs b/crates/merry-runtime/src/agent_loop/types.rs index ac376e04..d68dd216 100644 --- a/crates/merry-runtime/src/agent_loop/types.rs +++ b/crates/merry-runtime/src/agent_loop/types.rs @@ -57,7 +57,7 @@ impl AgentLoopResult { &self.status } - /// Runtime events collected in emission order. + /// Retained execution evidence in emission order, excluding live text deltas and rate samples. #[must_use] pub fn events(&self) -> &[RuntimeJournalEvent] { &self.events @@ -87,7 +87,7 @@ impl AgentLoopResult { self.session_usage.as_ref() } - /// Consumes the result and returns the collected events. + /// Consumes the result and returns retained evidence, excluding transient observations. #[must_use] pub fn into_events(self) -> Vec { self.events @@ -160,7 +160,7 @@ impl AgentLoopError { } } - /// Runtime events collected before the method error. + /// Retained evidence before the method error, excluding live deltas and rate samples. #[must_use] pub fn events(&self) -> &[RuntimeJournalEvent] { &self.events diff --git a/crates/merry-runtime/src/compaction/window.rs b/crates/merry-runtime/src/compaction/window.rs index 6ef724de..c710d164 100644 --- a/crates/merry-runtime/src/compaction/window.rs +++ b/crates/merry-runtime/src/compaction/window.rs @@ -2,6 +2,7 @@ use super::CompactionError; use crate::{ checkpoint::CheckpointRef, session::{ModelTurnId, ModelTurnStatus}, + token_estimate::TokenEstimateScale, }; use merry_core::{ArtifactId, ToolCallId, ToolCallResultStatus}; use serde::Serialize; @@ -16,6 +17,7 @@ pub(crate) struct CompactionWindowBudget { replacement_fixed_dynamic_body_tokens: u64, archive_only_fixed_dynamic_body_tokens: u64, checkpoint_output_ceiling_tokens: u64, + token_estimate_scale: TokenEstimateScale, } impl CompactionWindowBudget { @@ -46,6 +48,7 @@ impl CompactionWindowBudget { replacement_fixed_dynamic_body_tokens, archive_only_fixed_dynamic_body_tokens, checkpoint_output_ceiling_tokens, + token_estimate_scale: TokenEstimateScale::default(), }) } @@ -56,7 +59,20 @@ impl CompactionWindowBudget { Self::new(u64::MAX, u64::MAX, 0, 0, checkpoint_output_ceiling_tokens) } - /// Adds a bounded raw-history target to fixed input and the summary ceiling. + /// Applies the primary request calibration to destination-body estimates only. + pub(crate) fn with_token_estimate_scale(self, scale: TokenEstimateScale) -> Self { + Self { + token_estimate_scale: scale, + ..self + } + } + + /// Converts fixed context, checkpoint, and retained history to calibrated tokens. + pub(crate) fn estimate_body_tokens(self, base_tokens: u64) -> u64 { + self.token_estimate_scale.estimate(base_tokens) + } + + /// Adds a bounded history target to calibrated fixed input and summary ceiling. /// The hard body budget remains authoritative; arithmetic overflow is rejected. pub(crate) fn with_retained_history_target( self, @@ -65,6 +81,7 @@ impl CompactionWindowBudget { let preferred_tokens = self .replacement_fixed_dynamic_body_tokens .checked_add(self.checkpoint_output_ceiling_tokens) + .map(|tokens| self.estimate_body_tokens(tokens)) .and_then(|tokens| tokens.checked_add(retained_history_tokens)) .ok_or(CompactionError::BudgetOverflow)?; Ok(Self { diff --git a/crates/merry-runtime/src/events/journal_stream.rs b/crates/merry-runtime/src/events/journal_stream.rs index 5d2cd559..6286f162 100644 --- a/crates/merry-runtime/src/events/journal_stream.rs +++ b/crates/merry-runtime/src/events/journal_stream.rs @@ -17,13 +17,24 @@ use std::{ }, task::{Context, Poll}, }; +use tokio::sync::watch; use tokio::task::JoinHandle; use tokio_stream::wrappers::ReceiverStream; +use tokio_stream::wrappers::WatchStream; use tokio_util::sync::CancellationToken; -/// One atomic enqueue item in the internal journal channel. +/// One atomic enqueue item in the internal semantic journal channel. pub(crate) struct RuntimeJournalEventBatch(RuntimeJournalEventBatchKind); +pub(crate) type RuntimeRateUpdateSender = watch::Sender>; + +pub(crate) fn runtime_rate_update_channel() -> ( + RuntimeRateUpdateSender, + watch::Receiver>, +) { + watch::channel(None) +} + // The single-event path is hot; boxing it would add an allocation to every journal emission. #[allow(clippy::large_enum_variant)] enum RuntimeJournalEventBatchKind { @@ -58,6 +69,15 @@ impl RuntimeJournalEventBatch { } } + fn first_sequence(&self) -> u64 { + match &self.0 { + RuntimeJournalEventBatchKind::Single(event) => event.sequence, + RuntimeJournalEventBatchKind::Multiple(events) => { + events.first().map_or(0, |event| event.sequence) + } + } + } + fn into_iter(self) -> RuntimeJournalEventBatchIter { match self.0 { RuntimeJournalEventBatchKind::Single(event) => { @@ -100,10 +120,17 @@ impl Iterator for RuntimeJournalEventBatchIter { /// when callers want the producer to finish normally. Dropping it is the /// cancellation path for the active step. The permit may remain active after /// the producer stops while an in-flight persistence transaction finishes -/// discarding staged state or installing a durable commit. +/// discarding staged state or installing a durable commit. Semantic events are +/// delivered in sequence order; replaceable output-rate observations use a +/// latest-only channel and intermediate samples may be skipped. pub struct RuntimeJournalEventStream { inner: Option>, pending: Option, + pending_sequence: Option, + pending_batch_started: bool, + rate_updates: WatchStream>, + pending_rate_update: Option, + rate_updates_closed: bool, cancellation_token: CancellationToken, producer_handle: Option>, } @@ -111,12 +138,18 @@ pub struct RuntimeJournalEventStream { impl RuntimeJournalEventStream { pub(crate) fn new( inner: ReceiverStream, + rate_updates: watch::Receiver>, cancellation_token: CancellationToken, producer_handle: JoinHandle<()>, ) -> Self { Self { inner: Some(inner), pending: None, + pending_sequence: None, + pending_batch_started: false, + rate_updates: WatchStream::from_changes(rate_updates), + pending_rate_update: None, + rate_updates_closed: false, cancellation_token, producer_handle: Some(producer_handle), } @@ -128,25 +161,62 @@ impl Stream for RuntimeJournalEventStream { fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { loop { - if let Some(event) = self.pending.as_mut().and_then(Iterator::next) { - return Poll::Ready(Some(event)); + if !self.rate_updates_closed { + match Pin::new(&mut self.rate_updates).poll_next(cx) { + Poll::Ready(Some(Some(event))) => self.pending_rate_update = Some(event), + Poll::Ready(Some(None)) => {} + Poll::Ready(None) => self.rate_updates_closed = true, + Poll::Pending => {} + } } - self.pending = None; - - let Some(inner) = self.inner.as_mut() else { - return Poll::Ready(None); - }; - match Pin::new(inner).poll_next(cx) { - Poll::Ready(Some(batch)) => { - self.pending = Some(batch.into_iter()); + let rate_precedes_pending_batch = !self.pending_batch_started + && self.pending_rate_update.as_ref().is_some_and(|rate| { + self.pending_sequence + .is_some_and(|sequence| rate.sequence < sequence) + }); + if let Some(pending) = self.pending.as_mut() { + if rate_precedes_pending_batch { + return Poll::Ready(self.pending_rate_update.take()); } - Poll::Ready(None) => { - self.producer_handle.take(); - return Poll::Ready(None); + let event = pending.next(); + if event.is_some() { + self.pending_batch_started = true; + return Poll::Ready(event); } - Poll::Pending => return Poll::Pending, + self.pending = None; + self.pending_sequence = None; + self.pending_batch_started = false; } + + if let Some(inner) = self.inner.as_mut() { + match Pin::new(inner).poll_next(cx) { + Poll::Ready(Some(batch)) => { + self.pending_sequence = Some(batch.first_sequence()); + self.pending = Some(batch.into_iter()); + self.pending_batch_started = false; + continue; + } + Poll::Ready(None) => { + self.inner = None; + } + Poll::Pending => { + if let Some(rate) = self.pending_rate_update.take() { + return Poll::Ready(Some(rate)); + } + } + } + } + + if let Some(rate) = self.pending_rate_update.take() { + return Poll::Ready(Some(rate)); + } + + if self.inner.is_none() && self.rate_updates_closed { + self.producer_handle.take(); + return Poll::Ready(None); + } + return Poll::Pending; } } } @@ -199,7 +269,9 @@ struct ActiveStepPermitInner { #[cfg(test)] mod tests { use super::*; + use futures_util::StreamExt; use merry_core::{RuntimeJournalPayload, SessionId}; + use tokio::sync::mpsc; fn event(sequence: u64) -> RuntimeJournalEvent { RuntimeJournalEvent::new( @@ -209,6 +281,14 @@ mod tests { ) } + fn rate_event(sequence: u64) -> RuntimeJournalEvent { + RuntimeJournalEvent::new( + SessionId::new("journal-event-batch-test").expect("valid session id"), + sequence, + RuntimeJournalPayload::ModelOutputRateUpdated { rate: None }, + ) + } + #[test] fn event_batch_rejects_empty_and_preserves_all_events_in_order() { assert!(RuntimeJournalEventBatch::from_events(Vec::new()).is_none()); @@ -223,4 +303,78 @@ mod tests { [4, 5, 6] ); } + + #[tokio::test] + async fn rate_updates_replace_unread_values_instead_of_accumulating() { + let (rate_sender, rate_receiver) = runtime_rate_update_channel(); + let (event_sender, event_receiver) = mpsc::channel(1); + event_sender + .send( + RuntimeJournalEventBatch::from_events(vec![event(1)]) + .expect("non-empty event batch"), + ) + .await + .expect("event stream should be open"); + rate_sender.send_replace(Some(rate_event(2))); + rate_sender.send_replace(Some(rate_event(3))); + drop(rate_sender); + drop(event_sender); + + let producer_handle = tokio::spawn(async {}); + let mut stream = RuntimeJournalEventStream::new( + ReceiverStream::new(event_receiver), + rate_receiver, + CancellationToken::new(), + producer_handle, + ); + let events: Vec<_> = stream.by_ref().collect().await; + + assert_eq!( + events + .iter() + .map(|event| event.sequence) + .collect::>(), + vec![1, 3] + ); + assert!(matches!( + events[1].payload, + RuntimeJournalPayload::ModelOutputRateUpdated { rate: None } + )); + } + + #[tokio::test] + async fn a_new_rate_replaces_a_pending_rate_before_rendering_it() { + let (rate_sender, rate_receiver) = runtime_rate_update_channel(); + let (event_sender, event_receiver) = mpsc::channel(1); + event_sender + .send( + RuntimeJournalEventBatch::from_events(vec![event(1), event(2)]) + .expect("non-empty event batch"), + ) + .await + .expect("event stream should be open"); + drop(event_sender); + rate_sender.send_replace(Some(rate_event(3))); + + let producer_handle = tokio::spawn(async {}); + let mut stream = RuntimeJournalEventStream::new( + ReceiverStream::new(event_receiver), + rate_receiver, + CancellationToken::new(), + producer_handle, + ); + + assert_eq!( + stream.next().await.expect("first semantic event").sequence, + 1 + ); + rate_sender.send_replace(Some(rate_event(4))); + assert_eq!( + stream.next().await.expect("second semantic event").sequence, + 2 + ); + assert_eq!(stream.next().await.expect("latest rate event").sequence, 4); + drop(rate_sender); + assert!(stream.next().await.is_none()); + } } diff --git a/crates/merry-runtime/src/events/mod.rs b/crates/merry-runtime/src/events/mod.rs index 20cf749e..a25e3983 100644 --- a/crates/merry-runtime/src/events/mod.rs +++ b/crates/merry-runtime/src/events/mod.rs @@ -6,6 +6,9 @@ mod public_stream; mod tool_output; pub use journal_stream::RuntimeJournalEventStream; -pub(crate) use journal_stream::{ActiveStepPermit, RuntimeJournalEventBatch}; +pub(crate) use journal_stream::{ + ActiveStepPermit, RuntimeJournalEventBatch, RuntimeRateUpdateSender, + runtime_rate_update_channel, +}; pub use projector::RuntimeEventProjector; pub use public_stream::RuntimeEventStream; diff --git a/crates/merry-runtime/src/events/projector.rs b/crates/merry-runtime/src/events/projector.rs index 03a7dbad..d8f08c76 100644 --- a/crates/merry-runtime/src/events/projector.rs +++ b/crates/merry-runtime/src/events/projector.rs @@ -35,6 +35,9 @@ impl RuntimeEventProjector { Ok(Some(RuntimeEvent::SessionStarted { source })) } RuntimeJournalPayload::StepStarted => Ok(Some(RuntimeEvent::StepStarted { source })), + RuntimeJournalPayload::ModelOutputRateUpdated { rate } => { + Ok(Some(RuntimeEvent::ModelOutputRateUpdated { rate, source })) + } RuntimeJournalPayload::StepCompleted => { Ok(Some(RuntimeEvent::StepCompleted { source })) } diff --git a/crates/merry-runtime/src/events/public_stream.rs b/crates/merry-runtime/src/events/public_stream.rs index 6679e941..5c7bedae 100644 --- a/crates/merry-runtime/src/events/public_stream.rs +++ b/crates/merry-runtime/src/events/public_stream.rs @@ -9,16 +9,23 @@ use std::{ pin::Pin, task::{Context, Poll}, }; +use tokio::sync::watch; use tokio::{sync::mpsc, task::JoinHandle}; use tokio_stream::wrappers::ReceiverStream; +use tokio_stream::wrappers::WatchStream; /// Stream of SDK-facing runtime events. /// /// Dropping this stream aborts its projection task. The projection task owns the /// underlying journal stream, so aborting it drops the journal stream and -/// preserves existing runtime-step cancellation behavior. +/// preserves existing runtime-step cancellation behavior. Semantic events are +/// buffered in order; output-rate events are latest-only telemetry and may be +/// coalesced while the consumer is busy. pub struct RuntimeEventStream { inner: Option>, + rate_updates: WatchStream>, + rate_updates_closed: bool, + pending_rate_update: Option, producer_handle: Option>, } @@ -29,12 +36,16 @@ impl RuntimeEventStream { buffer_size: usize, ) -> Self { let (sender, receiver) = mpsc::channel(buffer_size); + let (rate_sender, rate_receiver) = watch::channel(None); let producer_handle = tokio::spawn(async move { - project_journal_stream(journal_stream, runtime, sender).await; + project_journal_stream(journal_stream, runtime, sender, rate_sender).await; }); Self { inner: Some(ReceiverStream::new(receiver)), + rate_updates: WatchStream::from_changes(rate_receiver), + rate_updates_closed: false, + pending_rate_update: None, producer_handle: Some(producer_handle), } } @@ -44,17 +55,35 @@ impl Stream for RuntimeEventStream { type Item = RuntimeEvent; fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - let Some(inner) = self.inner.as_mut() else { - return Poll::Ready(None); - }; + if !self.rate_updates_closed { + match Pin::new(&mut self.rate_updates).poll_next(cx) { + Poll::Ready(Some(Some(event))) => self.pending_rate_update = Some(event), + Poll::Ready(Some(None)) => {} + Poll::Ready(None) => self.rate_updates_closed = true, + Poll::Pending => {} + } + } - match Pin::new(inner).poll_next(cx) { - Poll::Ready(None) => { - self.producer_handle.take(); - Poll::Ready(None) + if let Some(inner) = self.inner.as_mut() { + match Pin::new(inner).poll_next(cx) { + Poll::Ready(Some(event)) => return Poll::Ready(Some(event)), + Poll::Ready(None) => self.inner = None, + Poll::Pending => { + if let Some(event) = self.pending_rate_update.take() { + return Poll::Ready(Some(event)); + } + } } - poll => poll, } + + if let Some(event) = self.pending_rate_update.take() { + return Poll::Ready(Some(event)); + } + if self.inner.is_none() && self.rate_updates_closed { + self.producer_handle.take(); + return Poll::Ready(None); + } + Poll::Pending } } @@ -72,6 +101,7 @@ async fn project_journal_stream( mut journal_stream: RuntimeJournalEventStream, runtime: Runtime, sender: mpsc::Sender, + rate_sender: watch::Sender>, ) { let mut projector = RuntimeEventProjector::new(); @@ -82,7 +112,9 @@ async fn project_journal_stream( continue; }; - if sender.send(public_event).await.is_err() { + if matches!(&public_event, RuntimeEvent::ModelOutputRateUpdated { .. }) { + rate_sender.send_replace(Some(public_event)); + } else if sender.send(public_event).await.is_err() { break; } } diff --git a/crates/merry-runtime/src/interactive.rs b/crates/merry-runtime/src/interactive.rs index c3e78521..ac53413d 100644 --- a/crates/merry-runtime/src/interactive.rs +++ b/crates/merry-runtime/src/interactive.rs @@ -14,7 +14,7 @@ mod types; use crate::{AgentLoopConfig, Runtime, RuntimeError, StepContext}; use producer::{InteractiveProducer, InteractiveProducerInput}; use std::sync::{Arc, atomic::AtomicU64}; -use tokio::sync::mpsc; +use tokio::sync::{mpsc, watch}; use tokio_stream::wrappers::ReceiverStream; pub use self::handles::{ @@ -42,6 +42,7 @@ impl Runtime { let subagent_completion_notify = self.subagent_completion_notify(); let run_id = types::next_interactive_run_id(); let (message_sender, message_receiver) = mpsc::channel(16); + let (rate_sender, rate_receiver) = watch::channel(None); let (command_sender, command_receiver) = mpsc::channel(16); let (bridge_sender, bridge_receiver) = mpsc::channel(4); let bridge_resolution_epoch = Arc::new(AtomicU64::new(0)); @@ -52,6 +53,7 @@ impl Runtime { plan_event_receiver, subagent_completion_notify, message_sender, + rate_sender, bridge_receiver, bridge_resolution_epoch: Arc::clone(&bridge_resolution_epoch), loop_token: producer_token, @@ -65,6 +67,7 @@ impl Runtime { InteractiveRunEventStream::new( run_id, ReceiverStream::new(message_receiver), + rate_receiver, loop_token, producer_handle, bridge_sender, diff --git a/crates/merry-runtime/src/interactive/handles.rs b/crates/merry-runtime/src/interactive/handles.rs index 780ce5a1..f0f67516 100644 --- a/crates/merry-runtime/src/interactive/handles.rs +++ b/crates/merry-runtime/src/interactive/handles.rs @@ -13,7 +13,7 @@ use std::sync::{ atomic::{AtomicU64, Ordering}, }; use tokio::{ - sync::{mpsc, oneshot}, + sync::{mpsc, oneshot, watch}, task::JoinHandle, }; use tokio_stream::wrappers::ReceiverStream; @@ -31,6 +31,9 @@ pub use crate::agent_loop::AgentRunMessage as InteractiveRunMessage; /// bindings. pub struct InteractiveRunEventStream { inner: Option>, + inner_closed: bool, + rate_updates: watch::Receiver>, + rate_updates_closed: bool, cancellation_token: CancellationToken, producer_handle: Option>>, bridge_sender: mpsc::Sender, @@ -46,6 +49,7 @@ impl InteractiveRunEventStream { pub(super) fn new( run_id: InteractiveRunId, inner: ReceiverStream, + rate_updates: watch::Receiver>, cancellation_token: CancellationToken, producer_handle: JoinHandle>, bridge_sender: mpsc::Sender, @@ -54,6 +58,9 @@ impl InteractiveRunEventStream { let observed_bridge_resolution_epoch = bridge_resolution_epoch.load(Ordering::Acquire); Self { inner: Some(inner), + inner_closed: false, + rate_updates, + rate_updates_closed: false, cancellation_token, producer_handle: Some(producer_handle), bridge_sender, @@ -106,9 +113,51 @@ impl InteractiveRunEventStream { return Some(message); } use futures_util::StreamExt; - match self.inner.as_mut() { - Some(inner) => inner.next().await, - None => None, + loop { + if self.inner_closed && self.rate_updates_closed { + return None; + } + if self.inner_closed { + match self.rate_updates.changed().await { + Ok(()) => { + if let Some(event) = self.rate_updates.borrow_and_update().clone() { + return Some(InteractiveRunMessage::Event(event)); + } + } + Err(_) => self.rate_updates_closed = true, + } + continue; + } + let Some(inner) = self.inner.as_mut() else { + self.inner_closed = true; + continue; + }; + if self.rate_updates_closed { + let message = inner.next().await; + if message.is_none() { + self.inner_closed = true; + } + return message; + } + tokio::select! { + biased; + message = inner.next() => { + if message.is_none() { + self.inner_closed = true; + } + return message; + } + changed = self.rate_updates.changed() => { + match changed { + Ok(()) => { + if let Some(event) = self.rate_updates.borrow_and_update().clone() { + return Some(InteractiveRunMessage::Event(event)); + } + } + Err(_) => self.rate_updates_closed = true, + } + } + } } } diff --git a/crates/merry-runtime/src/interactive/producer.rs b/crates/merry-runtime/src/interactive/producer.rs index fe161125..a367c89c 100644 --- a/crates/merry-runtime/src/interactive/producer.rs +++ b/crates/merry-runtime/src/interactive/producer.rs @@ -14,7 +14,7 @@ use crate::{AgentLoopConfig, Runtime, bridge::BridgeToolResultCommand, events::A use merry_llm::GenerationConfig; use std::collections::BTreeSet; use std::sync::{Arc, atomic::AtomicU64}; -use tokio::sync::{Notify, mpsc}; +use tokio::sync::{Notify, mpsc, watch}; use tokio_util::sync::CancellationToken; pub(super) struct InteractiveProducer { @@ -24,6 +24,7 @@ pub(super) struct InteractiveProducer { pub(super) plan_event_receiver: crate::plan::PlanControllerEventReceiver, pub(super) subagent_completion_notify: Option>, pub(super) message_sender: mpsc::Sender, + pub(super) rate_sender: watch::Sender>, pub(super) bridge_receiver: mpsc::Receiver, pub(super) bridge_resolution_epoch: Arc, pub(super) bridge_pending: bool, @@ -50,6 +51,7 @@ pub(super) struct InteractiveProducerInput { pub(super) plan_event_receiver: crate::plan::PlanControllerEventReceiver, pub(super) subagent_completion_notify: Option>, pub(super) message_sender: mpsc::Sender, + pub(super) rate_sender: watch::Sender>, pub(super) bridge_receiver: mpsc::Receiver, pub(super) bridge_resolution_epoch: Arc, pub(super) loop_token: CancellationToken, @@ -66,6 +68,7 @@ impl InteractiveProducer { plan_event_receiver, subagent_completion_notify, message_sender, + rate_sender, bridge_receiver, bridge_resolution_epoch, loop_token, @@ -80,6 +83,7 @@ impl InteractiveProducer { plan_event_receiver, subagent_completion_notify, message_sender, + rate_sender, bridge_receiver, bridge_resolution_epoch, bridge_pending: false, diff --git a/crates/merry-runtime/src/interactive/producer/model.rs b/crates/merry-runtime/src/interactive/producer/model.rs index d60fbc87..06c0442f 100644 --- a/crates/merry-runtime/src/interactive/producer/model.rs +++ b/crates/merry-runtime/src/interactive/producer/model.rs @@ -77,7 +77,9 @@ impl InteractiveProducer { let Some(event) = event else { return Some(events); }; - events.push(event.clone()); + if !event.payload.is_transient() { + events.push(event.clone()); + } if !self.project_and_send_runtime_event(&mut projector, event).await { return None; diff --git a/crates/merry-runtime/src/interactive/producer/output.rs b/crates/merry-runtime/src/interactive/producer/output.rs index fab6a1e5..7292c391 100644 --- a/crates/merry-runtime/src/interactive/producer/output.rs +++ b/crates/merry-runtime/src/interactive/producer/output.rs @@ -65,6 +65,13 @@ impl InteractiveProducer { } pub(super) async fn send_event(&self, event: RuntimeEvent) -> bool { + if matches!(&event, RuntimeEvent::ModelOutputRateUpdated { .. }) { + if self.rate_sender.is_closed() { + return false; + } + self.rate_sender.send_replace(Some(event)); + return true; + } self.message_sender .send(InteractiveRunMessage::Event(event)) .await diff --git a/crates/merry-runtime/src/model_completion.rs b/crates/merry-runtime/src/model_completion.rs index f84b0b38..55eb9fbb 100644 --- a/crates/merry-runtime/src/model_completion.rs +++ b/crates/merry-runtime/src/model_completion.rs @@ -105,7 +105,11 @@ pub(crate) async fn complete_single_text( }; match item { - Some(Ok(ModelEvent::Started | ModelEvent::OutputTextDelta { .. })) => {} + Some(Ok( + ModelEvent::Started + | ModelEvent::OutputTextDelta { .. } + | ModelEvent::OutputProgress { .. }, + )) => {} Some(Ok(ModelEvent::ToolCallRequested { .. })) => { return Err(ModelCompletionError::ToolCallRequested); } diff --git a/crates/merry-runtime/src/runtime.rs b/crates/merry-runtime/src/runtime.rs index c6a008db..fe961184 100644 --- a/crates/merry-runtime/src/runtime.rs +++ b/crates/merry-runtime/src/runtime.rs @@ -39,6 +39,7 @@ mod journal_persistence; mod memory_activation; mod model_output; mod model_turn_lifecycle; +mod output_rate; mod permission_execution; mod plan_read; mod plan_tool_execution; diff --git a/crates/merry-runtime/src/runtime/auto_compaction/mod.rs b/crates/merry-runtime/src/runtime/auto_compaction/mod.rs index d4d165a7..4e41c0e2 100644 --- a/crates/merry-runtime/src/runtime/auto_compaction/mod.rs +++ b/crates/merry-runtime/src/runtime/auto_compaction/mod.rs @@ -164,6 +164,7 @@ impl CompactionRequestBudget { fixed_dynamic_body_tokens.archive_only, checkpoint_output_ceiling_tokens, )? + .with_token_estimate_scale(request_budget.token_estimate_scale) .with_retained_history_target(resolved_budget.retained_history_token_target())?; Ok(Self { source, diff --git a/crates/merry-runtime/src/runtime/auto_compaction/phase.rs b/crates/merry-runtime/src/runtime/auto_compaction/phase.rs index 426e1293..f26eb453 100644 --- a/crates/merry-runtime/src/runtime/auto_compaction/phase.rs +++ b/crates/merry-runtime/src/runtime/auto_compaction/phase.rs @@ -24,7 +24,7 @@ use super::{ use crate::{ CitationCompactionPolicy, CompactionError, CompactionOutcome, compaction::{ArchiveOnlyCompactionInput, CompactionPreparation}, - events::{ActiveStepPermit, RuntimeJournalEventBatch}, + events::{ActiveStepPermit, RuntimeJournalEventBatch, RuntimeRateUpdateSender}, step::StepInput, }; use merry_core::{ErrorInfo, ToolSpec}; @@ -72,6 +72,7 @@ pub(in crate::runtime) enum HardWatermarkOutcome { pub(in crate::runtime) async fn reduce_context_at_hard_watermark( inner: &Arc, sender: &mpsc::Sender, + rate_sender: &RuntimeRateUpdateSender, token: &CancellationToken, active_permit: &ActiveStepPermit, parts: HardWatermarkCompaction<'_>, @@ -193,7 +194,7 @@ pub(in crate::runtime) async fn reduce_context_at_hard_watermark( } } CompactionAttempt::Generate(plan) => { - if !send_compaction_started_event(inner, sender, token).await { + if !send_compaction_started_event(inner, sender, rate_sender, token).await { return HardWatermarkOutcome::Aborted; } match generate_and_install_compaction( diff --git a/crates/merry-runtime/src/runtime/auto_compaction/source.rs b/crates/merry-runtime/src/runtime/auto_compaction/source.rs index eccdec37..3d343e3a 100644 --- a/crates/merry-runtime/src/runtime/auto_compaction/source.rs +++ b/crates/merry-runtime/src/runtime/auto_compaction/source.rs @@ -64,6 +64,11 @@ pub(super) async fn manual_compaction_budget( config.provider().capabilities(), &request, context_window_override, + inner + .session + .lock() + .await + .token_estimate_scale(config.provider().name(), &request), )?; let source = CompactionRequestSource::new(request, &history_ids, 0).map_err(|error| { RuntimeError::CompactionModelRequest { diff --git a/crates/merry-runtime/src/runtime/journal_emission.rs b/crates/merry-runtime/src/runtime/journal_emission.rs index 0b9ba498..173c88d5 100644 --- a/crates/merry-runtime/src/runtime/journal_emission.rs +++ b/crates/merry-runtime/src/runtime/journal_emission.rs @@ -8,6 +8,7 @@ use merry_core::{ CompactionUsageWindow, ErrorInfo, ModelUsage, PendingToolCall, RuntimeJournalEvent, RuntimeJournalPayload, ToolCallResultStatus, UsageContextWindow, }; +use tokio::sync::watch; use tokio::sync::{mpsc, mpsc::Permit}; use tokio_util::sync::CancellationToken; @@ -80,7 +81,7 @@ pub(super) async fn send_assistant_text_output_delta_event( session.record_transient_event(RuntimeJournalPayload::AssistantOutputDelta { delta }) }; - inner.emit_journal_batch(permit, event.into()); + inner.emit_journal_batch(permit, event.clone().into()); true } @@ -194,6 +195,7 @@ pub(super) async fn send_model_usage_updated_event( model_usage: ModelUsage, context: Option, compaction: Option, + observation: Option, ) -> Result { if token.is_cancelled() { return Ok(false); @@ -208,16 +210,62 @@ pub(super) async fn send_model_usage_updated_event( if token.is_cancelled() { return Ok(false); } - session.record_model_usage(model_usage, context, compaction)? + let event = session.record_model_usage(model_usage, context, compaction)?; + if let Some(observation) = observation { + session.calibrate_request_tokens(observation, model_usage); + } + event }; - inner.emit_journal_batch(permit, event.into()); + inner.emit_journal_batch(permit, event.clone().into()); Ok(true) } +/// Whether contention may discard an intermediate sample or must preserve a boundary. +pub(super) enum RateEventDelivery { + Replaceable, + Boundary, +} + +/// Sends a boundary or best-effort sample through the latest-only rate channel; +/// false means cancellation or receiver closure. Intermediate replaceable +/// samples never wait behind the semantic event queue. +pub(super) async fn send_model_output_rate_event( + inner: &RuntimeInner, + rate_sender: &watch::Sender>, + token: &CancellationToken, + rate: Option, + delivery: RateEventDelivery, +) -> bool { + if token.is_cancelled() { + return false; + } + if rate_sender.is_closed() { + return matches!(delivery, RateEventDelivery::Replaceable); + } + let event = { + let mut session = match delivery { + RateEventDelivery::Replaceable => { + let Ok(session) = inner.session.try_lock() else { + return true; + }; + session + } + RateEventDelivery::Boundary => inner.session.lock().await, + }; + if token.is_cancelled() { + return false; + } + session.record_transient_event(RuntimeJournalPayload::ModelOutputRateUpdated { rate }) + }; + rate_sender.send_replace(Some(event)); + true +} + pub(super) async fn send_compaction_started_event( inner: &RuntimeInner, sender: &mpsc::Sender, + rate_sender: &watch::Sender>, token: &CancellationToken, ) -> bool { if token.is_cancelled() { @@ -228,15 +276,18 @@ pub(super) async fn send_compaction_started_event( return false; }; - let event = { + let (event, reset) = { let mut session = inner.session.lock().await; if token.is_cancelled() { return false; } - session.record_compaction_started() + let reset = session + .record_transient_event(RuntimeJournalPayload::ModelOutputRateUpdated { rate: None }); + (session.record_compaction_started(), reset) }; inner.emit_journal_batch(permit, event.into()); + rate_sender.send_replace(Some(reset)); true } @@ -410,6 +461,7 @@ pub(super) async fn send_normal_event( pub(super) async fn send_step_started_event( inner: &RuntimeInner, sender: &mpsc::Sender, + rate_sender: &watch::Sender>, token: &CancellationToken, ) -> Option { if token.is_cancelled() { @@ -417,14 +469,29 @@ pub(super) async fn send_step_started_event( } let permit = reserve_normal_event_slot(sender, token).await?; - let event = { + let (event, reset) = { let mut session = inner.session.lock().await; if token.is_cancelled() { return None; } - session.record_step_started() + let event = session.record_step_started(); + let reset = + if inner + .model_configs + .contains_role(crate::RuntimeModelRole::Primary) + { + Some(session.record_transient_event( + RuntimeJournalPayload::ModelOutputRateUpdated { rate: None }, + )) + } else { + None + }; + (event, reset) }; inner.emit_journal_batch(permit, event.clone().into()); + if let Some(reset) = reset { + rate_sender.send_replace(Some(reset)); + } Some(event) } diff --git a/crates/merry-runtime/src/runtime/output_rate.rs b/crates/merry-runtime/src/runtime/output_rate.rs new file mode 100644 index 00000000..4acdb390 --- /dev/null +++ b/crates/merry-runtime/src/runtime/output_rate.rs @@ -0,0 +1,51 @@ +//! Per-attempt throughput reduction, independent of event delivery and presentation. + +use crate::token_estimate::estimate_utf8_tokens; +use merry_core::{ModelOutputRate, ModelUsage, OutputTimingQuality}; +use merry_llm::ModelOutputProgress; + +#[derive(Default)] +pub(super) struct OutputRateTracker { + progress: Option, +} + +impl OutputRateTracker { + pub(super) fn observe( + &mut self, + progress: Option, + ) -> Option { + self.progress = progress; + self.rate(None) + } + + pub(super) fn rate(&self, usage: Option) -> Option { + let progress = self.progress?; + let tokens = usage.map_or_else( + || estimate_utf8_tokens(progress.utf8_bytes()), + |usage| usage.output_tokens(), + ); + let quality = if progress.timing_quality() == OutputTimingQuality::ReceiveWindow + && !progress.reasoning_observed() + && usage.is_some_and(|usage| { + usage + .reasoning_output_tokens() + .is_some_and(|tokens| tokens > 0) + }) { + OutputTimingQuality::PartialOutput + } else { + progress.timing_quality() + }; + Some( + ModelOutputRate::new( + tokens, + progress.elapsed(), + if usage.is_none() { + merry_core::OutputTokenSource::Estimated + } else { + merry_core::OutputTokenSource::ProviderUsage + }, + ) + .with_timing_quality(quality), + ) + } +} diff --git a/crates/merry-runtime/src/runtime/provider_request.rs b/crates/merry-runtime/src/runtime/provider_request.rs index 06d600c2..810a64f8 100644 --- a/crates/merry-runtime/src/runtime/provider_request.rs +++ b/crates/merry-runtime/src/runtime/provider_request.rs @@ -8,7 +8,7 @@ use crate::{ resolve_context_window, session::{SessionState, TranscriptItemSnapshot}, step::{StepInput, StepModelRequestParts, compile_step_model_request}, - token_estimate::estimate_model_input_tokens, + token_estimate::{TokenEstimateScale, estimate_model_input_tokens}, }; use merry_core::{CompactionUsageWindow, ErrorInfo, UsageContextWindow}; use merry_llm::{GenerationConfig, ModelError, ModelName}; @@ -248,6 +248,7 @@ pub(super) struct RequestContextBudget { pub(super) policy: ContextBudgetPolicy, pub(super) budget: ContextBudget, pub(super) dynamic_body_estimated_tokens: u64, + pub(super) token_estimate_scale: TokenEstimateScale, pub(super) decision: CheckpointDecision, } @@ -294,6 +295,7 @@ pub(super) fn request_context_budget( capabilities: &merry_llm::ModelCapabilities, request: &merry_llm::ModelRequest, context_window_override: Option, + token_estimate_scale: TokenEstimateScale, ) -> Result { let window = resolve_request_context_window(capabilities, context_window_override)?; let output_reserve_tokens = request @@ -309,11 +311,12 @@ pub(super) fn request_context_budget( let budget = ContextBudget::from_window( window.tokens(), DEFAULT_EFFECTIVE_CONTEXT_WINDOW_PERCENT, - stable_prefix_estimated_tokens, + token_estimate_scale.estimate(stable_prefix_estimated_tokens), output_reserve_tokens, policy, )?; - let dynamic_body_estimated_tokens = estimate_model_input_tokens(request.dynamic_input()); + let dynamic_body_estimated_tokens = + token_estimate_scale.estimate(estimate_model_input_tokens(request.dynamic_input())); let decision = decide_checkpoint(dynamic_body_estimated_tokens, budget); Ok(RequestContextBudget { @@ -321,6 +324,7 @@ pub(super) fn request_context_budget( policy, budget, dynamic_body_estimated_tokens, + token_estimate_scale, decision, }) } diff --git a/crates/merry-runtime/src/runtime/provider_step.rs b/crates/merry-runtime/src/runtime/provider_step.rs index a92082a0..15fc0809 100644 --- a/crates/merry-runtime/src/runtime/provider_step.rs +++ b/crates/merry-runtime/src/runtime/provider_step.rs @@ -1,44 +1,37 @@ +mod consume; + use super::auto_compaction::{ CompactionProgress, HardWatermarkCompaction, HardWatermarkOutcome, reduce_context_at_hard_watermark, }; use super::journal_emission::{ - send_assistant_text_output_completed_events, send_assistant_text_output_delta_event, send_cancelled_event, send_compaction_completed_event, send_failed_event, - send_model_tool_call_response_events, send_model_usage_updated_event, trace_provider_step_cancelled, trace_provider_step_failed, }; use super::memory_activation::{ ActivationProjectionGuard, clear_current_activated_memories, memory_activation_seed_from_step_input, }; -use super::model_output::{ - DIAGNOSTIC_MODEL_TOOL_CALL_MIXED_OUTPUT, diagnostic_from_model_error, is_cancelled_model_error, - pending_tool_call_from_model, pending_tool_calls_from_outputs, record_streamed_tool_call, - tool_call_commentary_text, -}; +use super::model_output::{diagnostic_from_model_error, is_cancelled_model_error}; use super::model_turn_lifecycle::{InProgressModelTurnGuard, cancel_model_turn, fail_model_turn}; use super::provider_request::{ compile_step_request_from_inputs, request_context_budget, step_request_compile_diagnostic, step_request_inputs_from_session, step_usage_context_snapshot, trace_provider_request, trace_provider_request_budget_unavailable, }; -use super::provider_stream::{ - stream_model_with_retry_policy, wait_for_model_stream_item, wait_for_retrying_stream_setup, -}; +use super::provider_stream::{stream_model_with_retry_policy, wait_for_retrying_stream_setup}; use super::{DIAGNOSTIC_TOOL_CALL_RESULT_REQUIRED, RuntimeInner, diagnostic_from_text}; use crate::{ CheckpointDecision, - events::{ActiveStepPermit, RuntimeJournalEventBatch}, + events::{ActiveStepPermit, RuntimeJournalEventBatch, RuntimeRateUpdateSender}, memory::MemoryActivationContext, model_config::ModelProviderConfig, plan::unix_time_ms, step::StepInput, }; -use merry_core::PendingToolCall; -use merry_llm::{FinishReason, GenerationConfig, ModelEvent, ModelOutput, ModelStreamContext}; +use merry_llm::{GenerationConfig, ModelStreamContext}; use std::sync::{Arc, atomic::Ordering}; use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; @@ -51,6 +44,7 @@ async fn has_unresolved_pending_tool_calls(inner: &RuntimeInner) -> bool { pub(super) struct ProviderStepControl<'a> { token: &'a CancellationToken, active_permit: &'a ActiveStepPermit, + rate_sender: &'a RuntimeRateUpdateSender, step_sequence: u64, } @@ -58,11 +52,13 @@ impl<'a> ProviderStepControl<'a> { pub(super) const fn new( token: &'a CancellationToken, active_permit: &'a ActiveStepPermit, + rate_sender: &'a RuntimeRateUpdateSender, step_sequence: u64, ) -> Self { Self { token, active_permit, + rate_sender, step_sequence, } } @@ -80,6 +76,7 @@ pub(super) async fn run_provider_step( let ProviderStepControl { token, active_permit, + rate_sender, step_sequence, } = control; if has_unresolved_pending_tool_calls(inner).await { @@ -303,8 +300,17 @@ pub(super) async fn run_provider_step( .read() .await .map(std::num::NonZeroU64::get); - let mut request_budget = - request_context_budget(provider.capabilities(), &request, context_window_override); + let mut token_estimate_scale = inner + .session + .lock() + .await + .token_estimate_scale(provider.name(), &request); + let mut request_budget = request_context_budget( + provider.capabilities(), + &request, + context_window_override, + token_estimate_scale, + ); // Read the compaction policy once for this step instead of re-locking per // compaction attempt. let (automatic_compaction_enabled, automatic_policy, compaction_reasoning_effort) = { @@ -353,6 +359,7 @@ pub(super) async fn run_provider_step( let outcome = reduce_context_at_hard_watermark( inner, sender, + rate_sender, token, active_permit, HardWatermarkCompaction { @@ -414,8 +421,17 @@ pub(super) async fn run_provider_step( return; } }; - request_budget = - request_context_budget(provider.capabilities(), &request, context_window_override); + token_estimate_scale = inner + .session + .lock() + .await + .token_estimate_scale(provider.name(), &request); + request_budget = request_context_budget( + provider.capabilities(), + &request, + context_window_override, + token_estimate_scale, + ); current_budget = match &request_budget { Ok(budget) => *budget, Err(error) => { @@ -478,6 +494,8 @@ pub(super) async fn run_provider_step( .observe_model_request(&request, step_sequence); let usage_context_snapshot = step_usage_context_snapshot(request_budget.as_ref().ok(), automatic_compaction_enabled); + let request_token_observation = + crate::token_estimate::RequestTokenObservation::new(provider.name(), &request); let sent_continuation_count = request.continuations().len(); let turn_id = { @@ -559,7 +577,7 @@ pub(super) async fn run_provider_step( } }; - let mut stream = match stream_result { + let stream = match stream_result { Ok(stream) => { tracing::debug!( category = "provider_setup_success", @@ -587,210 +605,19 @@ pub(super) async fn run_provider_step( }; projection_guard.disarm(); - let mut commentary_text = String::new(); - let mut streamed_tool_calls: Vec = Vec::new(); - - loop { - let item = wait_for_model_stream_item( - inner, - sender, - token, - &mut stream, - &mut retry_event_receiver, - ) - .await; - - let item = match item { - Some(item) => item, - None => { - cancel_model_turn(inner, sender, turn_id).await; - return; - } - }; - - match item { - Some(Ok(ModelEvent::Started)) => { - tracing::debug!(category = "started", "runtime model stream event received"); - } - Some(Ok(ModelEvent::OutputTextDelta { delta })) => { - if !delta.is_empty() { - tracing::trace!( - category = "output_text_delta_nonempty", - "runtime model stream event received" - ); - commentary_text.push_str(&delta); - if !send_assistant_text_output_delta_event(inner, sender, token, delta).await { - cancel_model_turn(inner, sender, turn_id).await; - return; - } - } - } - Some(Ok(ModelEvent::Completed { response })) => { - tracing::debug!( - category = "completed", - finish_reason = ?response.finish_reason(), - "runtime model stream event received" - ); - if let Some(model_usage) = response.usage() { - match send_model_usage_updated_event( - inner, - sender, - token, - model_usage, - usage_context_snapshot.context, - usage_context_snapshot.compaction, - ) - .await - { - Ok(true) => {} - Ok(false) => { - cancel_model_turn(inner, sender, turn_id).await; - return; - } - Err(diagnostic) => { - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - } - } - } - match response.finish_reason() { - FinishReason::Stop => { - if !streamed_tool_calls.is_empty() { - let diagnostic = diagnostic_from_text( - DIAGNOSTIC_MODEL_TOOL_CALL_MIXED_OUTPUT, - "model requested a tool call before completing with text output", - ); - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - } - - let [ModelOutput::Text { text }] = response.outputs() else { - let diagnostic = diagnostic_from_text( - "model_output_unsupported", - "model stop output must contain exactly one text item", - ); - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - }; - - if !send_assistant_text_output_completed_events( - inner, - sender, - token, - turn_id, - text.clone(), - ) - .await - { - cancel_model_turn(inner, sender, turn_id).await; - } - return; - } - FinishReason::ToolCalls => { - match pending_tool_calls_from_outputs( - response.outputs(), - &streamed_tool_calls, - ) { - Ok(calls) => { - if calls.len() > 1 - && final_output_contract.as_ref().is_some_and(|contract| { - calls.iter().any(|call| call.name() == contract.tool_name()) - }) - { - let diagnostic = diagnostic_from_text( - "final_output_tool_batch_mixed", - "final-output tool calls must be the only call in their model batch", - ); - fail_model_turn(inner, sender, token, turn_id, diagnostic) - .await; - return; - } - let commentary = - tool_call_commentary_text(response.outputs(), &commentary_text); - let sent = send_model_tool_call_response_events( - inner, sender, token, turn_id, commentary, calls, - ) - .await; - if !sent { - cancel_model_turn(inner, sender, turn_id).await; - } - } - Err(diagnostic) => { - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - } - } - return; - } - FinishReason::Length => { - let diagnostic = diagnostic_from_text( - "model_length", - "model output stopped because it reached a length limit", - ); - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - } - FinishReason::Blocked => { - let diagnostic = diagnostic_from_text( - "model_blocked", - "model output was blocked by provider safety or content policy", - ); - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - } - FinishReason::Cancelled => { - cancel_model_turn(inner, sender, turn_id).await; - return; - } - FinishReason::Error => { - let diagnostic = diagnostic_from_text( - "model_finish_error", - "model output stopped because the provider reported a finish error", - ); - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - } - } - } - Some(Ok(ModelEvent::ToolCallRequested { call })) => { - tracing::debug!( - category = "tool_call_requested", - "runtime model stream event received" - ); - match pending_tool_call_from_model(&call) - .and_then(|call| record_streamed_tool_call(&mut streamed_tool_calls, call)) - { - Ok(()) => {} - Err(diagnostic) => { - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - } - } - } - Some(Err(error)) => { - let error_kind = error.kind(); - tracing::debug!( - category = "provider_error", - error_kind = ?error_kind, - "runtime model stream event received" - ); - if is_cancelled_model_error(&error) { - cancel_model_turn(inner, sender, turn_id).await; - return; - } - - let diagnostic = diagnostic_from_model_error(error); - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - } - None => { - tracing::debug!(category = "eof", "runtime model stream ended"); - let diagnostic = diagnostic_from_text( - "model_stream_eof", - "model stream ended before completion", - ); - fail_model_turn(inner, sender, token, turn_id, diagnostic).await; - return; - } - } - } + consume::consume_model_stream( + inner, + sender, + rate_sender, + token, + consume::ModelStreamRun { + stream, + retry_event_receiver, + turn_id, + usage_context_snapshot, + request_token_observation, + final_output_contract, + }, + ) + .await; } diff --git a/crates/merry-runtime/src/runtime/provider_step/consume.rs b/crates/merry-runtime/src/runtime/provider_step/consume.rs new file mode 100644 index 00000000..93f24ef5 --- /dev/null +++ b/crates/merry-runtime/src/runtime/provider_step/consume.rs @@ -0,0 +1,288 @@ +//! Consumes one model turn, reducing streamed output and committing its terminal result. + +use super::super::{ + RuntimeInner, diagnostic_from_text, + journal_emission::{ + RateEventDelivery, send_assistant_text_output_completed_events, + send_assistant_text_output_delta_event, send_model_output_rate_event, + send_model_tool_call_response_events, send_model_usage_updated_event, + }, + model_output::{ + DIAGNOSTIC_MODEL_TOOL_CALL_MIXED_OUTPUT, diagnostic_from_model_error, + is_cancelled_model_error, pending_tool_call_from_model, pending_tool_calls_from_outputs, + record_streamed_tool_call, tool_call_commentary_text, + }, + model_turn_lifecycle::{cancel_model_turn, fail_model_turn}, + output_rate::OutputRateTracker, + provider_request::StepUsageContextSnapshot, + provider_stream::wait_for_model_stream_item, +}; +use crate::{ + events::{RuntimeJournalEventBatch, RuntimeRateUpdateSender}, + session::ModelTurnId, + token_estimate::RequestTokenObservation, +}; +use merry_core::PendingToolCall; +use merry_llm::{FinishReason, ModelEvent, ModelEventStream, ModelOutput, ModelRetryEvent}; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +pub(super) struct ModelStreamRun { + pub(super) stream: ModelEventStream, + pub(super) retry_event_receiver: mpsc::Receiver, + pub(super) turn_id: ModelTurnId, + pub(super) usage_context_snapshot: StepUsageContextSnapshot, + pub(super) request_token_observation: RequestTokenObservation, + pub(super) final_output_contract: Option, +} + +pub(super) async fn consume_model_stream( + inner: &RuntimeInner, + sender: &mpsc::Sender, + rate_sender: &RuntimeRateUpdateSender, + token: &CancellationToken, + run: ModelStreamRun, +) { + let ModelStreamRun { + mut stream, + mut retry_event_receiver, + turn_id, + usage_context_snapshot, + request_token_observation, + final_output_contract, + } = run; + let mut commentary_text = String::new(); + let mut output_rate = OutputRateTracker::default(); + let mut streamed_tool_calls: Vec = Vec::new(); + + loop { + let item = wait_for_model_stream_item( + inner, + sender, + token, + &mut stream, + &mut retry_event_receiver, + ) + .await; + + let item = match item { + Some(item) => item, + None => { + cancel_model_turn(inner, sender, turn_id).await; + return; + } + }; + + match item { + Some(Ok(ModelEvent::OutputProgress { progress })) => { + let rate = output_rate.observe(progress); + let delivery = if rate.is_some() { + RateEventDelivery::Replaceable + } else { + RateEventDelivery::Boundary + }; + if !send_model_output_rate_event(inner, rate_sender, token, rate, delivery).await { + cancel_model_turn(inner, sender, turn_id).await; + return; + } + } + Some(Ok(ModelEvent::Started)) => { + tracing::debug!(category = "started", "runtime model stream event received"); + } + Some(Ok(ModelEvent::OutputTextDelta { delta })) => { + if !delta.is_empty() { + tracing::trace!( + category = "output_text_delta_nonempty", + "runtime model stream event received" + ); + commentary_text.push_str(&delta); + if !send_assistant_text_output_delta_event(inner, sender, token, delta).await { + cancel_model_turn(inner, sender, turn_id).await; + return; + } + } + } + Some(Ok(ModelEvent::Completed { response })) => { + if let Some(rate) = output_rate.rate(response.usage()) + && !send_model_output_rate_event( + inner, + rate_sender, + token, + Some(rate), + RateEventDelivery::Boundary, + ) + .await + { + cancel_model_turn(inner, sender, turn_id).await; + return; + } + tracing::debug!( + category = "completed", + finish_reason = ?response.finish_reason(), + "runtime model stream event received" + ); + if let Some(model_usage) = response.usage() { + match send_model_usage_updated_event( + inner, + sender, + token, + model_usage, + usage_context_snapshot.context, + usage_context_snapshot.compaction, + (response.finish_reason() != FinishReason::Cancelled) + .then_some(request_token_observation), + ) + .await + { + Ok(true) => {} + Ok(false) => { + cancel_model_turn(inner, sender, turn_id).await; + return; + } + Err(diagnostic) => { + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + } + } + } + match response.finish_reason() { + FinishReason::Stop => { + if !streamed_tool_calls.is_empty() { + let diagnostic = diagnostic_from_text( + DIAGNOSTIC_MODEL_TOOL_CALL_MIXED_OUTPUT, + "model requested a tool call before completing with text output", + ); + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + } + + let [ModelOutput::Text { text }] = response.outputs() else { + let diagnostic = diagnostic_from_text( + "model_output_unsupported", + "model stop output must contain exactly one text item", + ); + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + }; + + if !send_assistant_text_output_completed_events( + inner, + sender, + token, + turn_id, + text.clone(), + ) + .await + { + cancel_model_turn(inner, sender, turn_id).await; + } + return; + } + FinishReason::ToolCalls => { + match pending_tool_calls_from_outputs( + response.outputs(), + &streamed_tool_calls, + ) { + Ok(calls) => { + if calls.len() > 1 + && final_output_contract.as_ref().is_some_and(|contract| { + calls.iter().any(|call| call.name() == contract.tool_name()) + }) + { + let diagnostic = diagnostic_from_text( + "final_output_tool_batch_mixed", + "final-output tool calls must be the only call in their model batch", + ); + fail_model_turn(inner, sender, token, turn_id, diagnostic) + .await; + return; + } + let commentary = + tool_call_commentary_text(response.outputs(), &commentary_text); + let sent = send_model_tool_call_response_events( + inner, sender, token, turn_id, commentary, calls, + ) + .await; + if !sent { + cancel_model_turn(inner, sender, turn_id).await; + } + } + Err(diagnostic) => { + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + } + } + return; + } + FinishReason::Length => { + let diagnostic = diagnostic_from_text( + "model_length", + "model output stopped because it reached a length limit", + ); + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + } + FinishReason::Blocked => { + let diagnostic = diagnostic_from_text( + "model_blocked", + "model output was blocked by provider safety or content policy", + ); + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + } + FinishReason::Cancelled => { + cancel_model_turn(inner, sender, turn_id).await; + return; + } + FinishReason::Error => { + let diagnostic = diagnostic_from_text( + "model_finish_error", + "model output stopped because the provider reported a finish error", + ); + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + } + } + } + Some(Ok(ModelEvent::ToolCallRequested { call })) => { + tracing::debug!( + category = "tool_call_requested", + "runtime model stream event received" + ); + match pending_tool_call_from_model(&call) + .and_then(|call| record_streamed_tool_call(&mut streamed_tool_calls, call)) + { + Ok(()) => {} + Err(diagnostic) => { + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + } + } + } + Some(Err(error)) => { + let error_kind = error.kind(); + tracing::debug!( + category = "provider_error", + error_kind = ?error_kind, + "runtime model stream event received" + ); + if is_cancelled_model_error(&error) { + cancel_model_turn(inner, sender, turn_id).await; + return; + } + + let diagnostic = diagnostic_from_model_error(error); + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + } + None => { + tracing::debug!(category = "eof", "runtime model stream ended"); + let diagnostic = diagnostic_from_text( + "model_stream_eof", + "model stream ended before completion", + ); + fail_model_turn(inner, sender, token, turn_id, diagnostic).await; + return; + } + } + } +} diff --git a/crates/merry-runtime/src/runtime/step.rs b/crates/merry-runtime/src/runtime/step.rs index 1f5c8ce5..89002bf3 100644 --- a/crates/merry-runtime/src/runtime/step.rs +++ b/crates/merry-runtime/src/runtime/step.rs @@ -1,7 +1,10 @@ use super::{Runtime, RuntimeInner, journal_emission::*, provider_step}; use crate::{ RuntimeError, RuntimeModelRole, - events::{ActiveStepPermit, RuntimeJournalEventBatch, RuntimeJournalEventStream}, + events::{ + ActiveStepPermit, RuntimeJournalEventBatch, RuntimeJournalEventStream, + runtime_rate_update_channel, + }, step::{StepContext, StepInput}, }; use merry_llm::GenerationConfig; @@ -11,6 +14,11 @@ use tokio_stream::wrappers::ReceiverStream; use tokio_util::sync::CancellationToken; use tracing::Instrument; +struct StepEventSenders { + journal: mpsc::Sender, + rate: crate::events::RuntimeRateUpdateSender, +} + impl Runtime { pub(crate) fn step_with_active_permit( &self, @@ -22,6 +30,7 @@ impl Runtime { let step_token = parent_token.child_token(); let producer_token = step_token.clone(); let (sender, receiver) = mpsc::channel(self.inner.event_buffer_size.get()); + let (rate_sender, rate_receiver) = runtime_rate_update_channel(); let inner = Arc::clone(&self.inner); let producer_span = tracing::debug_span!( "runtime.step", @@ -39,7 +48,10 @@ impl Runtime { async move { run_step( inner, - sender, + StepEventSenders { + journal: sender, + rate: rate_sender, + }, producer_token, input, generation_config, @@ -53,6 +65,7 @@ impl Runtime { Ok(RuntimeJournalEventStream::new( ReceiverStream::new(receiver), + rate_receiver, step_token, producer_handle, )) @@ -61,13 +74,17 @@ impl Runtime { async fn run_step( inner: Arc, - sender: mpsc::Sender, + event_senders: StepEventSenders, token: CancellationToken, input: StepInput, generation_config: GenerationConfig, final_output_contract: Option, active_permit: ActiveStepPermit, ) { + let StepEventSenders { + journal: sender, + rate, + } = event_senders; tracing::debug!(category = "started", "runtime step started"); if token.is_cancelled() { @@ -94,7 +111,7 @@ async fn run_step( return; } - let Some(step_started) = send_step_started_event(&inner, &sender, &token).await else { + let Some(step_started) = send_step_started_event(&inner, &sender, &rate, &token).await else { let _ = send_cancelled_if_requested(&inner, &sender, &token).await; return; }; @@ -126,7 +143,12 @@ async fn run_step( provider_step::run_provider_step( &inner, &sender, - provider_step::ProviderStepControl::new(&token, &active_permit, step_started.sequence), + provider_step::ProviderStepControl::new( + &token, + &active_permit, + &rate, + step_started.sequence, + ), input, generation_config, final_output_contract, diff --git a/crates/merry-runtime/src/runtime/tests/bridge_tool_flow.rs b/crates/merry-runtime/src/runtime/tests/bridge_tool_flow.rs index fa30a3be..69dd7e99 100644 --- a/crates/merry-runtime/src/runtime/tests/bridge_tool_flow.rs +++ b/crates/merry-runtime/src/runtime/tests/bridge_tool_flow.rs @@ -260,8 +260,12 @@ async fn invalid_bridge_terminal_events_share_one_slot_without_stranding_produce .sequence, 1 ); + assert!(matches!( + first_stream.next().await.expect("rate reset event").payload, + RuntimeJournalPayload::ModelOutputRateUpdated { rate: None } + )); let pending = first_stream.next().await.expect("pending tool event"); - assert_eq!(pending.sequence, 2); + assert_eq!(pending.sequence, 3); assert!(matches!( pending.payload, RuntimeJournalPayload::ToolCallPending { .. } @@ -290,6 +294,6 @@ async fn invalid_bridge_terminal_events_share_one_slot_without_stranding_produce let second_events = second_stream.collect::>().await; assert_eq!( second_events.first().expect("second step event").sequence, - 5 + 6 ); } diff --git a/crates/merry-runtime/src/runtime/tests/compaction_transaction.rs b/crates/merry-runtime/src/runtime/tests/compaction_transaction.rs index e0f85d33..cafd3e75 100644 --- a/crates/merry-runtime/src/runtime/tests/compaction_transaction.rs +++ b/crates/merry-runtime/src/runtime/tests/compaction_transaction.rs @@ -437,7 +437,9 @@ async fn automatic_compaction_completed_waits_for_directory_durability() { loop { let event = events.next().await.expect("compaction should start"); match event.payload { - RuntimeJournalPayload::SessionStarted | RuntimeJournalPayload::StepStarted => {} + RuntimeJournalPayload::SessionStarted + | RuntimeJournalPayload::StepStarted + | RuntimeJournalPayload::ModelOutputRateUpdated { rate: None } => {} RuntimeJournalPayload::CompactionStarted => break, payload => panic!("unexpected event before compaction starts: {payload:?}"), } diff --git a/crates/merry-runtime/src/runtime/tests/context_cache.rs b/crates/merry-runtime/src/runtime/tests/context_cache.rs index 18df1213..76cb4a80 100644 --- a/crates/merry-runtime/src/runtime/tests/context_cache.rs +++ b/crates/merry-runtime/src/runtime/tests/context_cache.rs @@ -175,8 +175,13 @@ async fn soft_watermark_does_not_call_the_compaction_provider() { let requests = primary.recorded_requests(); assert_eq!(requests.len(), 1); - let budget = request_context_budget(primary.capabilities(), &requests[0], None) - .expect("request budget should resolve"); + let budget = request_context_budget( + primary.capabilities(), + &requests[0], + None, + Default::default(), + ) + .expect("request budget should resolve"); assert_eq!(budget.decision, CheckpointDecision::PlanCheckpoint); assert!( events diff --git a/crates/merry-runtime/src/runtime/tests/model_role_flow.rs b/crates/merry-runtime/src/runtime/tests/model_role_flow.rs index 3e828a71..71eb232c 100644 --- a/crates/merry-runtime/src/runtime/tests/model_role_flow.rs +++ b/crates/merry-runtime/src/runtime/tests/model_role_flow.rs @@ -38,6 +38,7 @@ mod automatic_compaction; mod budget; mod rolling_compaction; +mod token_calibration; mod manual_compaction; diff --git a/crates/merry-runtime/src/runtime/tests/model_role_flow/automatic_compaction.rs b/crates/merry-runtime/src/runtime/tests/model_role_flow/automatic_compaction.rs index 5fad34a9..e0fa186c 100644 --- a/crates/merry-runtime/src/runtime/tests/model_role_flow/automatic_compaction.rs +++ b/crates/merry-runtime/src/runtime/tests/model_role_flow/automatic_compaction.rs @@ -110,13 +110,13 @@ async fn hard_watermark_auto_compaction_emits_lifecycle_events() { "StepCompleted" ] ); - assert!(matches!( - events[2].payload, + assert!(events.iter().any(|event| matches!( + event.payload, RuntimeJournalPayload::CompactionCompleted { ref checkpoint_id, covered_history_item_count: 2 } if checkpoint_id.starts_with("checkpoint-auto-compaction-events-") - )); + ))); // The output ceiling is the checkpoint text budget plus the reasoning // reserve sized from the request input, so it must exceed the text budget // while the whole request still fits the primary window. diff --git a/crates/merry-runtime/src/runtime/tests/model_role_flow/budget.rs b/crates/merry-runtime/src/runtime/tests/model_role_flow/budget.rs index 470ca7df..a50a943b 100644 --- a/crates/merry-runtime/src/runtime/tests/model_role_flow/budget.rs +++ b/crates/merry-runtime/src/runtime/tests/model_role_flow/budget.rs @@ -33,8 +33,8 @@ fn request_context_budget_uses_dynamic_estimate_watermarks() { ) .expect("valid request"); - let budget = - request_context_budget(&capabilities, &request, None).expect("budget should calculate"); + let budget = request_context_budget(&capabilities, &request, None, Default::default()) + .expect("budget should calculate"); assert_eq!( budget.window.source(), @@ -55,6 +55,70 @@ fn request_context_budget_uses_dynamic_estimate_watermarks() { ); } +#[test] +fn calibrated_budget_counts_fixed_input_and_body_without_scaling_output_reserve() { + let capabilities = ModelCapabilities::new(true, true, false, true, Some(100_000), Some(10_000)) + .expect("capabilities"); + let provider = merry_core::ProviderName::new("calibrated-provider").expect("provider"); + let request = ModelRequest::new_with_continuations_and_stable_prefix( + named_model("calibrated-model"), + vec![ + ModelMessage::new( + ModelMessageRole::System, + ModelContent::text(&"rules ".repeat(8_000)).expect("text"), + ) + .expect("message"), + ModelMessage::new( + ModelMessageRole::User, + ModelContent::text(&"body ".repeat(40_000)).expect("text"), + ) + .expect("message"), + ], + Vec::new(), + Vec::new(), + GenerationConfig::default(), + 1, + ) + .expect("request"); + let calibration = crate::token_estimate::RequestTokenCalibration::observe( + None, + crate::token_estimate::RequestTokenObservation::new(&provider, &request), + crate::token_estimate::estimate_request_input_tokens(&request) * 2, + ) + .expect("calibration"); + let baseline = request_context_budget(&capabilities, &request, None, Default::default()) + .expect("baseline budget"); + let corrected = request_context_budget( + &capabilities, + &request, + None, + calibration.scale_for(&provider, &request), + ) + .expect("corrected budget"); + assert_eq!(baseline.decision, CheckpointDecision::Continue); + assert_eq!(corrected.decision, CheckpointDecision::RequireCheckpoint); + assert_eq!( + corrected.budget.stable_prefix_tokens(), + baseline.budget.stable_prefix_tokens() * 2 + ); + assert_eq!( + corrected.dynamic_body_estimated_tokens, + baseline.dynamic_body_estimated_tokens * 2 + ); + assert_eq!( + corrected.budget.effective_window_tokens(), + baseline.budget.effective_window_tokens() + ); + assert_eq!( + corrected.budget.output_reserve_tokens(), + baseline.budget.output_reserve_tokens() + ); + assert_eq!( + corrected.budget.hard_water_tokens(), + baseline.budget.hard_water_tokens() - baseline.budget.stable_prefix_tokens() + ); +} + #[test] fn request_context_budget_derives_default_output_reserve_from_window() { let request = ModelRequest::new_with_continuations_and_stable_prefix( @@ -84,8 +148,8 @@ fn request_context_budget_derives_default_output_reserve_from_window() { ] { let capabilities = ModelCapabilities::new(true, true, false, true, Some(window), None) .expect("valid capabilities"); - let budget = - request_context_budget(&capabilities, &request, None).expect("budget should calculate"); + let budget = request_context_budget(&capabilities, &request, None, Default::default()) + .expect("budget should calculate"); assert_eq!( budget.budget.output_reserve_tokens(), @@ -114,8 +178,8 @@ fn request_context_budget_uses_codex_style_fallback_for_unknown_models() { ) .expect("valid request"); - let budget = - request_context_budget(&capabilities, &request, None).expect("budget should calculate"); + let budget = request_context_budget(&capabilities, &request, None, Default::default()) + .expect("budget should calculate"); assert_eq!(budget.window.tokens(), 272_000); assert_eq!(budget.window.source(), crate::ContextWindowSource::Fallback); @@ -142,7 +206,7 @@ fn request_context_budget_prefers_an_explicit_window_override() { ) .expect("valid request"); - let budget = request_context_budget(&capabilities, &request, Some(128_000)) + let budget = request_context_budget(&capabilities, &request, Some(128_000), Default::default()) .expect("budget should calculate"); assert_eq!(budget.window.tokens(), 128_000); diff --git a/crates/merry-runtime/src/runtime/tests/model_role_flow/rolling_compaction.rs b/crates/merry-runtime/src/runtime/tests/model_role_flow/rolling_compaction.rs index dc51c5c9..ad0bd6a9 100644 --- a/crates/merry-runtime/src/runtime/tests/model_role_flow/rolling_compaction.rs +++ b/crates/merry-runtime/src/runtime/tests/model_role_flow/rolling_compaction.rs @@ -368,6 +368,7 @@ async fn assert_one_shot_reduction( primary.capabilities(), final_request, Some(window_tokens), + Default::default(), ) .expect("primary budget"); assert!( diff --git a/crates/merry-runtime/src/runtime/tests/model_role_flow/token_calibration.rs b/crates/merry-runtime/src/runtime/tests/model_role_flow/token_calibration.rs new file mode 100644 index 00000000..86599987 --- /dev/null +++ b/crates/merry-runtime/src/runtime/tests/model_role_flow/token_calibration.rs @@ -0,0 +1,359 @@ +use crate::{ + CitationCompactionPolicy, CompactionConfig, FileSessionStore, Runtime, RuntimeModelRole, + StepContext, + runtime::tests::support::{ + common::{collect_step, completed_event_with, model_name, named_model, session_id}, + model_provider::{RecordingModelProvider, ScriptedModelProviderResponse}, + }, + token_estimate::{estimate_model_input_tokens, estimate_request_input_tokens}, +}; +use merry_core::{ModelUsage, ProviderName, RuntimeJournalPayload}; +use merry_llm::{ + FinishReason, ModelCapabilities, ModelError, ModelEvent, ModelEventStream, ModelOutput, + ModelProvider, ModelProviderFuture, ModelRequest, ModelResponse, ModelStreamContext, + ProviderErrorKind, +}; +use std::{ + collections::VecDeque, + sync::{Arc, Mutex}, +}; + +#[derive(Clone, Copy, Default)] +enum Feedback { + #[default] + Measured, + Missing, + Zero, + Cancelled, + Failed, +} + +struct MeteredProvider { + name: ProviderName, + capabilities: ModelCapabilities, + requests: Mutex>, + feedback: Mutex>, +} + +impl MeteredProvider { + fn new(feedback: Vec) -> Self { + Self { + name: ProviderName::new("metered-provider").expect("provider name"), + capabilities: ModelCapabilities::new(true, true, false, true, Some(64_000), Some(512)) + .expect("capabilities"), + requests: Mutex::new(Vec::new()), + feedback: Mutex::new(feedback.into()), + } + } + + fn last_request(&self) -> ModelRequest { + self.requests + .lock() + .expect("requests") + .last() + .expect("request") + .clone() + } +} + +impl ModelProvider for MeteredProvider { + fn name(&self) -> &ProviderName { + &self.name + } + + fn capabilities(&self) -> &ModelCapabilities { + &self.capabilities + } + + fn stream_model<'a>( + &'a self, + request: ModelRequest, + context: ModelStreamContext, + ) -> ModelProviderFuture<'a, Result> { + Box::pin(async move { + if context.cancellation_token().is_cancelled() { + return Err(ModelError::Cancelled); + } + let base_tokens = estimate_request_input_tokens(&request); + self.requests.lock().expect("requests").push(request); + let feedback = self + .feedback + .lock() + .expect("feedback") + .pop_front() + .unwrap_or_default(); + if matches!(feedback, Feedback::Failed) { + return Err(ModelError::provider( + ProviderErrorKind::InvalidRequest, + "fixture failure", + )); + } + let actual_tokens = match feedback { + Feedback::Cancelled => base_tokens * 100, + Feedback::Zero => 0, + _ => base_tokens * 2, + }; + let usage = + (!matches!(feedback, Feedback::Missing)).then_some(ModelUsage::with_details( + actual_tokens, + Some(actual_tokens), + 1, + None, + actual_tokens + 1, + )); + let finish = if matches!(feedback, Feedback::Cancelled) { + FinishReason::Cancelled + } else { + FinishReason::Stop + }; + let stream: ModelEventStream = + Box::pin(futures_util::stream::iter([Ok(ModelEvent::Completed { + response: ModelResponse::new(vec![ModelOutput::text("done")], finish, usage), + })])); + Ok(stream) + }) + } +} + +fn compactor() -> RecordingModelProvider { + let candidate = r#"{ + "confirmed_decisions": [], "rejected_approaches": [], + "constraints_preferences_boundaries": [], "corrected_misunderstandings": [], + "durable_conclusions": [{"id":"c1", "text":"Earlier history was compacted.", "refs":["h0"]}], + "open_questions": [], "current_progress_and_next_steps": [], "exact_details": [], "handoffs": [] + }"#; + RecordingModelProvider::with_script_and_capabilities( + (0..4) + .map(|_| { + ScriptedModelProviderResponse::Stream(vec![Ok(completed_event_with( + vec![ModelOutput::text(candidate)], + FinishReason::Stop, + ))]) + }) + .collect(), + ModelCapabilities::new(true, true, false, true, Some(256_000), None).expect("capabilities"), + ) +} + +fn policy() -> CitationCompactionPolicy { + CitationCompactionPolicy::new(Some(512), Some(16_384), 1).expect("policy") +} + +async fn assert_prediction(runtime: &Runtime, provider: &MeteredProvider, multiplier: u64) { + let usage = runtime.usage().await.expect("usage"); + let request = provider.last_request(); + assert_eq!( + usage.last.input_tokens(), + estimate_request_input_tokens(&request) * 2 + ); + assert_eq!( + usage.last.cached_input_tokens(), + Some(usage.last.input_tokens()) + ); + let compaction = usage.compaction.expect("compaction snapshot"); + assert_eq!( + compaction.dynamic_body_estimated_tokens, + Some(estimate_model_input_tokens(request.dynamic_input()) * multiplier), + ); +} + +#[tokio::test] +async fn multi_turn_feedback_triggers_compaction_before_the_uncalibrated_estimate_would() { + for feedback in [Feedback::Missing, Feedback::Measured] { + let primary = Arc::new(MeteredProvider::new(vec![feedback; 12])); + let compactor = compactor(); + let runtime = Runtime::builder(session_id("calibration-multi-turn")) + .model_provider(primary.clone(), model_name()) + .model_provider_for_role( + RuntimeModelRole::ContextCompaction, + Arc::new(compactor.clone()), + named_model("compactor"), + ) + .automatic_compaction(CompactionConfig::enabled(policy())) + .build() + .expect("runtime"); + let mut compacted = false; + for turn in 0..12 { + let text = format!("Turn {turn}: {}", "abcd".repeat(3_000)); + let events = collect_step(&runtime, &text, StepContext::default()).await; + assert!( + events + .iter() + .any(|event| matches!(event.payload, RuntimeJournalPayload::StepCompleted)), + "turn {turn}: {events:?}" + ); + compacted |= events.iter().any(|event| { + matches!( + event.payload, + RuntimeJournalPayload::CompactionCompleted { .. } + ) + }); + if matches!(feedback, Feedback::Measured) { + assert_prediction(&runtime, &primary, if turn == 0 { 1 } else { 2 }).await; + } + } + assert_eq!(compacted, matches!(feedback, Feedback::Measured)); + assert_eq!(compactor.recorded_requests().is_empty(), !compacted); + } +} + +#[tokio::test] +async fn calibrated_retention_planning_fits_the_tail_before_installing_a_checkpoint() { + for (feedback, expected_covered_items) in [(Feedback::Missing, 4), (Feedback::Measured, 6)] { + let primary = Arc::new(MeteredProvider::new(vec![feedback; 4])); + let compactor = compactor(); + let runtime = Runtime::builder(session_id("calibrated-retention")) + .model_provider(primary, model_name()) + .model_provider_for_role( + RuntimeModelRole::ContextCompaction, + Arc::new(compactor.clone()), + named_model("compactor"), + ) + .automatic_compaction(CompactionConfig::disabled()) + .build() + .expect("runtime"); + for turn in 0..4 { + let events = collect_step( + &runtime, + &format!("Turn {turn}: {}", "abcd".repeat(3_000)), + StepContext::default(), + ) + .await; + assert!( + events + .iter() + .any(|event| matches!(event.payload, RuntimeJournalPayload::StepCompleted)) + ); + } + let outcome = runtime + .compact_context_once( + policy() + .with_retained_model_turns(2) + .expect("retention policy"), + StepContext::default(), + ) + .await + .expect("manual compaction") + .expect("checkpoint"); + assert_eq!(outcome.covered_history_item_count(), expected_covered_items); + assert_eq!(compactor.recorded_requests().len(), 1); + } +} + +#[tokio::test] +async fn manual_compaction_and_resumed_requests_keep_the_primary_calibration() { + let primary = Arc::new(MeteredProvider::new(Vec::new())); + let compactor = compactor(); + let runtime = Runtime::builder(session_id("calibration-manual-resume")) + .model_provider(primary.clone(), model_name()) + .model_provider_for_role( + RuntimeModelRole::ContextCompaction, + Arc::new(compactor.clone()), + named_model("compactor"), + ) + .automatic_compaction(CompactionConfig::disabled()) + .build() + .expect("runtime"); + for turn in 0..8 { + let events = collect_step( + &runtime, + &format!("Turn {turn}: {}", "abcd".repeat(3_000)), + StepContext::default(), + ) + .await; + assert!( + events + .iter() + .any(|event| matches!(event.payload, RuntimeJournalPayload::StepCompleted)) + ); + } + runtime + .compact_context_once(policy(), StepContext::default()) + .await + .expect("manual compaction") + .expect("checkpoint"); + assert_eq!( + compactor.recorded_requests().len(), + 1, + "retained history must be fitted in calibrated units" + ); + let directory = tempfile::tempdir().expect("store directory"); + let store = FileSessionStore::new(directory.path()); + runtime.save_session_to(store.clone()).await.expect("save"); + let resumed = Runtime::builder(runtime.session_id().clone()) + .model_provider(primary.clone(), model_name()) + .automatic_compaction(CompactionConfig::disabled()) + .resume_from_store(store.clone()) + .await + .expect("resume"); + let events = collect_step( + &resumed, + "Continue after checkpoint and resume.", + StepContext::default(), + ) + .await; + assert!( + events + .iter() + .any(|event| matches!(event.payload, RuntimeJournalPayload::StepCompleted)) + ); + assert_prediction(&resumed, &primary, 2).await; + let usage = resumed + .usage() + .await + .expect("usage") + .compaction + .expect("budget"); + assert!(usage.dynamic_body_estimated_tokens.expect("estimate") < usage.hard_water_tokens); + + let switched = Runtime::builder(runtime.session_id().clone()) + .model_provider(primary.clone(), named_model("different-model")) + .automatic_compaction(CompactionConfig::disabled()) + .resume_from_store(store) + .await + .expect("resume with another model"); + let events = collect_step( + &switched, + "New model must start without inherited feedback.", + StepContext::default(), + ) + .await; + assert!( + events + .iter() + .any(|event| matches!(event.payload, RuntimeJournalPayload::StepCompleted)) + ); + assert_prediction(&switched, &primary, 1).await; +} + +#[tokio::test] +async fn missing_zero_cancelled_and_failed_usage_do_not_replace_valid_feedback() { + let primary = Arc::new(MeteredProvider::new(vec![ + Feedback::Measured, + Feedback::Missing, + Feedback::Zero, + Feedback::Cancelled, + Feedback::Failed, + Feedback::Measured, + ])); + let runtime = Runtime::builder(session_id("calibration-invalid-feedback")) + .model_provider(primary.clone(), model_name()) + .build() + .expect("runtime"); + for _ in 0..5 { + collect_step(&runtime, "Previous request.", StepContext::default()).await; + } + let events = collect_step( + &runtime, + "Use the last valid calibration.", + StepContext::default(), + ) + .await; + assert!( + events + .iter() + .any(|event| matches!(event.payload, RuntimeJournalPayload::StepCompleted)), + "{events:?}" + ); + assert_prediction(&runtime, &primary, 2).await; +} diff --git a/crates/merry-runtime/src/runtime/tests/provider_step_flow.rs b/crates/merry-runtime/src/runtime/tests/provider_step_flow.rs index 3dab9a34..8d474333 100644 --- a/crates/merry-runtime/src/runtime/tests/provider_step_flow.rs +++ b/crates/merry-runtime/src/runtime/tests/provider_step_flow.rs @@ -1,4 +1,5 @@ mod atomic_response; mod memory_lifecycle; +mod output_rate; mod request_projection; mod retry; diff --git a/crates/merry-runtime/src/runtime/tests/provider_step_flow/atomic_response.rs b/crates/merry-runtime/src/runtime/tests/provider_step_flow/atomic_response.rs index 778bbc0b..581f1a3a 100644 --- a/crates/merry-runtime/src/runtime/tests/provider_step_flow/atomic_response.rs +++ b/crates/merry-runtime/src/runtime/tests/provider_step_flow/atomic_response.rs @@ -50,6 +50,10 @@ async fn cancellation_after_atomic_tool_response_preserves_awaiting_turn() { stream.next().await.expect("step start event").payload, RuntimeJournalPayload::StepStarted )); + assert!(matches!( + stream.next().await.expect("rate reset event").payload, + RuntimeJournalPayload::ModelOutputRateUpdated { rate: None } + )); let commentary = stream.next().await.expect("commentary event"); assert!(matches!( commentary.payload, diff --git a/crates/merry-runtime/src/runtime/tests/provider_step_flow/output_rate.rs b/crates/merry-runtime/src/runtime/tests/provider_step_flow/output_rate.rs new file mode 100644 index 00000000..8e08c702 --- /dev/null +++ b/crates/merry-runtime/src/runtime/tests/provider_step_flow/output_rate.rs @@ -0,0 +1,265 @@ +use crate::{ + Runtime, StepContext, StepInput, + runtime::tests::support::{ + common::{model_name, session_id}, + model_provider::{RecordingModelProvider, ScriptedModelProviderResponse}, + }, +}; +use futures_util::StreamExt; +use merry_core::{ModelOutputRate, ProviderName, RuntimeEvent, RuntimeJournalPayload}; +use merry_llm::{ + FinishReason, ModelCapabilities, ModelError, ModelEvent, ModelEventStream, ModelOutput, + ModelOutputProgress, ModelProvider, ModelProviderFuture, ModelRequest, ModelResponse, + ModelStreamContext, Usage, +}; +use std::{num::NonZeroUsize, sync::Arc, time::Duration}; +use tokio::sync::Notify; + +struct CompletionNotifyingProvider { + inner: RecordingModelProvider, + completed: Arc, +} + +impl ModelProvider for CompletionNotifyingProvider { + fn name(&self) -> &ProviderName { + self.inner.name() + } + + fn capabilities(&self) -> &ModelCapabilities { + self.inner.capabilities() + } + + fn stream_model<'a>( + &'a self, + request: ModelRequest, + context: ModelStreamContext, + ) -> ModelProviderFuture<'a, Result> { + Box::pin(async move { + let completed = Arc::clone(&self.completed); + let stream = self.inner.stream_model(request, context).await?; + let stream: ModelEventStream = Box::pin(stream.inspect(move |event| { + if matches!(event, Ok(ModelEvent::Completed { .. })) { + completed.notify_one(); + } + })); + Ok(stream) + }) + } +} + +fn progress() -> ModelEvent { + ModelEvent::OutputProgress { + progress: Some( + ModelOutputProgress::new(1_600, Duration::from_secs(8)).with_reasoning_observed(true), + ), + } +} + +fn completed(usage: Option) -> ModelEvent { + ModelEvent::Completed { + response: ModelResponse::new(vec![ModelOutput::text("done")], FinishReason::Stop, usage), + } +} + +async fn rates(runtime: &Runtime) -> Vec> { + runtime + .stream( + StepInput::user_text("continue").unwrap(), + StepContext::default(), + ) + .unwrap() + .filter_map(|event| async move { + match event { + RuntimeEvent::ModelOutputRateUpdated { rate, .. } => Some(rate), + _ => None, + } + }) + .collect() + .await +} + +#[tokio::test] +async fn full_single_slot_buffer_does_not_block_replaceable_rate_updates() { + let completed = Arc::new(Notify::new()); + let events = (0..10_000) + .map(|_| Ok(progress())) + .chain([Ok(self::completed(Some(Usage::new(100, 2_400))))]) + .collect(); + let provider = CompletionNotifyingProvider { + inner: RecordingModelProvider::with_script(vec![ScriptedModelProviderResponse::Stream( + events, + )]), + completed: Arc::clone(&completed), + }; + let runtime = Runtime::builder(session_id("single-slot-rate-pressure")) + .model_provider(Arc::new(provider), model_name()) + .event_buffer_size(NonZeroUsize::new(1).unwrap()) + .build() + .unwrap(); + let mut stream = runtime + .step( + StepInput::user_text("continue").unwrap(), + StepContext::default(), + ) + .unwrap(); + assert!(matches!( + stream.next().await.unwrap().payload, + RuntimeJournalPayload::SessionStarted + )); + assert!(matches!( + stream.next().await.unwrap().payload, + RuntimeJournalPayload::StepStarted + )); + + tokio::time::timeout(Duration::from_secs(2), completed.notified()) + .await + .unwrap(); + + let events: Vec<_> = stream.collect().await; + let samples: Vec<_> = events + .iter() + .filter_map(|event| match event.payload { + RuntimeJournalPayload::ModelOutputRateUpdated { rate } => Some(rate), + _ => None, + }) + .collect(); + assert!( + samples.len() <= 3, + "replaceable rate updates must not accumulate a backlog: {samples:?}" + ); + assert_eq!(samples.last().unwrap().unwrap().output_tokens(), 2_400); + assert!(matches!( + events.last().unwrap().payload, + RuntimeJournalPayload::StepCompleted + )); +} + +#[tokio::test] +async fn output_rate_uses_complete_provider_usage_and_the_same_receive_interval() { + let usage = Usage::with_details(20_000, None, 2_400, Some(2_000), 22_400); + let provider = RecordingModelProvider::with_script(vec![ + ScriptedModelProviderResponse::Stream(vec![Ok(progress()), Ok(completed(Some(usage)))]), + ScriptedModelProviderResponse::Stream(vec![Ok(completed(Some(usage)))]), + ]); + let runtime = Runtime::builder(session_id("output-rate-usage")) + .model_provider(Arc::new(provider), model_name()) + .build() + .unwrap(); + let observed = rates(&runtime).await; + assert_eq!(observed.len(), 1, "rate snapshots are latest-only"); + assert_eq!( + observed[0], + Some(ModelOutputRate::new( + 2_400, + Duration::from_secs(8), + merry_core::OutputTokenSource::ProviderUsage + )) + ); + assert_eq!(observed[0].unwrap().tokens_per_second(), Some(300.0)); + assert_eq!( + rates(&runtime).await, + vec![None], + "an untimed next request cannot reuse the previous sample" + ); +} + +#[tokio::test] +async fn output_rate_keeps_estimates_when_usage_is_missing() { + let provider = RecordingModelProvider::with_script(vec![ + ScriptedModelProviderResponse::Stream(vec![Ok(progress()), Ok(completed(None))]), + ]); + let runtime = Runtime::builder(session_id("output-rate-no-usage")) + .model_provider(Arc::new(provider), model_name()) + .build() + .unwrap(); + let observed = rates(&runtime).await; + assert_eq!( + observed.last(), + Some(&Some(ModelOutputRate::new( + 400, + Duration::from_secs(8), + merry_core::OutputTokenSource::Estimated + ))) + ); +} + +#[tokio::test] +async fn cancelled_output_does_not_invent_a_completion_time_or_actual_usage() { + let provider = RecordingModelProvider::with_script(vec![ + ScriptedModelProviderResponse::Stream(vec![Ok(progress()), Err(ModelError::Cancelled)]), + ]); + let runtime = Runtime::builder(session_id("output-rate-cancelled")) + .model_provider(Arc::new(provider), model_name()) + .build() + .unwrap(); + assert_eq!( + rates(&runtime).await, + vec![Some(ModelOutputRate::new( + 400, + Duration::from_secs(8), + merry_core::OutputTokenSource::Estimated + ))] + ); +} + +#[tokio::test] +async fn retry_reset_cannot_pair_previous_attempt_timing_with_new_usage() { + let provider = + RecordingModelProvider::with_script(vec![ScriptedModelProviderResponse::Stream(vec![ + Ok(progress()), + Ok(ModelEvent::OutputProgress { progress: None }), + Ok(completed(Some(Usage::new(100, 500)))), + ])]); + let runtime = Runtime::builder(session_id("output-rate-reset")) + .model_provider(Arc::new(provider), model_name()) + .build() + .unwrap(); + assert_eq!(rates(&runtime).await, vec![None]); +} + +#[tokio::test] +async fn provider_usage_does_not_hide_partial_or_consumer_limited_timing() { + use merry_core::{OutputTimingQuality, OutputTokenSource}; + for quality in [ + OutputTimingQuality::ReceiveWindow, + OutputTimingQuality::PartialOutput, + OutputTimingQuality::ConsumerLimited, + ] { + let sample = + ModelOutputProgress::new(400, Duration::from_secs(2)).with_timing_quality(quality); + let provider = + RecordingModelProvider::with_script(vec![ScriptedModelProviderResponse::Stream(vec![ + Ok(ModelEvent::OutputProgress { + progress: Some(sample), + }), + Ok(completed(Some(Usage::with_details( + 10, + None, + 2_400, + Some(2_000), + 2_410, + )))), + ])]); + let runtime = Runtime::builder(session_id("output-rate-quality")) + .model_provider(Arc::new(provider), model_name()) + .build() + .unwrap(); + let observed = rates(&runtime).await; + let final_rate = observed.last().unwrap().unwrap(); + assert_eq!(final_rate.token_source(), OutputTokenSource::ProviderUsage); + assert_eq!(final_rate.output_tokens(), 2_400); + assert!(final_rate.is_estimated()); + assert_eq!(final_rate.tokens_per_second(), Some(1_200.0)); + if quality == OutputTimingQuality::ConsumerLimited { + assert_eq!( + final_rate.timing_quality(), + OutputTimingQuality::ConsumerLimited + ); + } else { + assert_eq!( + final_rate.timing_quality(), + OutputTimingQuality::PartialOutput + ); + } + } +} diff --git a/crates/merry-runtime/src/runtime/tests/provider_step_turn_lifecycle.rs b/crates/merry-runtime/src/runtime/tests/provider_step_turn_lifecycle.rs index 23aebbbc..574ed2e9 100644 --- a/crates/merry-runtime/src/runtime/tests/provider_step_turn_lifecycle.rs +++ b/crates/merry-runtime/src/runtime/tests/provider_step_turn_lifecycle.rs @@ -86,6 +86,12 @@ async fn one_slot_commentary_observes_atomically_committed_tool_response() { let step_started = stream.next().await.expect("step start event"); assert_eq!(session_started.sequence, 0); assert_eq!(step_started.sequence, 1); + let reset = stream.next().await.expect("rate reset event"); + assert_eq!(reset.sequence, 2); + assert!(matches!( + reset.payload, + RuntimeJournalPayload::ModelOutputRateUpdated { rate: None } + )); let session = loop { let session = runtime.inner.session.lock().await; @@ -108,14 +114,14 @@ async fn one_slot_commentary_observes_atomically_committed_tool_response() { }; let commentary = stream.next().await.expect("commentary event"); - assert_eq!(commentary.sequence, 2); + assert_eq!(commentary.sequence, 3); assert!(matches!( commentary.payload, RuntimeJournalPayload::AssistantOutputRecorded { .. } )); assert_eq!(session.pending_tool_calls().len(), 1); - assert_eq!(session.next_sequence(), 4); + assert_eq!(session.next_sequence(), 5); assert_eq!( session.model_turn_status(ModelTurnId::new(1)), Some(ModelTurnStatus::AwaitingToolResults) @@ -132,7 +138,7 @@ async fn one_slot_commentary_observes_atomically_committed_tool_response() { matches!( entry, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, kind: LedgerFactKind::ArtifactRecorded, .. } @@ -142,7 +148,7 @@ async fn one_slot_commentary_observes_atomically_committed_tool_response() { matches!( entry, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, kind: LedgerFactKind::ToolCallPending, .. } @@ -158,5 +164,5 @@ async fn one_slot_commentary_observes_atomically_committed_tool_response() { session.model_turn_status(ModelTurnId::new(1)), Some(ModelTurnStatus::AwaitingToolResults) ); - assert_eq!(session.next_sequence(), 4); + assert_eq!(session.next_sequence(), 5); } diff --git a/crates/merry-runtime/src/runtime/tests/rolling_compaction.rs b/crates/merry-runtime/src/runtime/tests/rolling_compaction.rs index c2da48e4..d3527b09 100644 --- a/crates/merry-runtime/src/runtime/tests/rolling_compaction.rs +++ b/crates/merry-runtime/src/runtime/tests/rolling_compaction.rs @@ -260,6 +260,7 @@ async fn run_three_cycle_case(window_tokens: u64) { &primary_capabilities(window_tokens), trigger_request, None, + Default::default(), ) .expect("request budget"); assert_eq!( diff --git a/crates/merry-runtime/src/runtime/tests/session_resume/savepoints.rs b/crates/merry-runtime/src/runtime/tests/session_resume/savepoints.rs index 4eff4aa3..5ad3f84c 100644 --- a/crates/merry-runtime/src/runtime/tests/session_resume/savepoints.rs +++ b/crates/merry-runtime/src/runtime/tests/session_resume/savepoints.rs @@ -159,6 +159,12 @@ async fn dropping_text_stream_while_savepoint_is_blocked_keeps_terminal_batch_co 0 ); assert_eq!(stream.next().await.expect("step start event").sequence, 1); + let reset = stream.next().await.expect("rate reset event"); + assert_eq!(reset.sequence, 2); + assert!(matches!( + reset.payload, + RuntimeJournalPayload::ModelOutputRateUpdated { rate: None } + )); let session = loop { let session = runtime.inner.session.lock().await; let assistant_recorded = session @@ -180,21 +186,21 @@ async fn dropping_text_stream_while_savepoint_is_blocked_keeps_terminal_batch_co }; let assistant = stream.next().await.expect("assistant output event"); - assert_eq!(assistant.sequence, 2); + assert_eq!(assistant.sequence, 3); assert!(matches!( assistant.payload, RuntimeJournalPayload::AssistantOutputRecorded { .. } )); assert_eq!( session.next_sequence(), - 4, + 5, "StepCompleted must commit before the terminal batch becomes observable" ); assert!(session.ledger_projection().entries().iter().any(|entry| { matches!( entry, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, kind: LedgerFactKind::StepCompleted, .. } @@ -202,7 +208,7 @@ async fn dropping_text_stream_while_savepoint_is_blocked_keeps_terminal_batch_co })); drop(stream); - assert_eq!(session.next_sequence(), 4); + assert_eq!(session.next_sequence(), 5); assert_eq!( session.model_turn_status(ModelTurnId::new(1)), Some(ModelTurnStatus::Completed) @@ -242,6 +248,10 @@ async fn terminal_journal_batch_waits_for_resume_savepoint_before_delivery() { events.next().await.expect("step start event").payload, RuntimeJournalPayload::StepStarted )); + assert!(matches!( + events.next().await.expect("rate reset event").payload, + RuntimeJournalPayload::ModelOutputRateUpdated { rate: None } + )); tokio::select! { biased; @@ -271,9 +281,9 @@ async fn terminal_journal_batch_waits_for_resume_savepoint_before_delivery() { .trajectory_snapshot() .await .expect("trajectory snapshot reads"); - assert_eq!(snapshot.latest_sequence(), 2); + assert_eq!(snapshot.latest_sequence(), 3); assert!(snapshot.records().iter().any(|record| { - record.start_sequence() == 2 + record.start_sequence() == 3 && record.status() == merry_core::TrajectoryRecordStatus::Succeeded })); } diff --git a/crates/merry-runtime/src/runtime/tests/support/common.rs b/crates/merry-runtime/src/runtime/tests/support/common.rs index d1b1cc6a..3c86fd5b 100644 --- a/crates/merry-runtime/src/runtime/tests/support/common.rs +++ b/crates/merry-runtime/src/runtime/tests/support/common.rs @@ -209,6 +209,12 @@ pub(in crate::runtime::tests) fn event_kind_names( ) -> Vec<&'static str> { events .iter() + .filter(|event| { + !matches!( + event.payload, + RuntimeJournalPayload::ModelOutputRateUpdated { .. } + ) + }) .map(|event| match event.payload { RuntimeJournalPayload::SessionStarted => "SessionStarted", RuntimeJournalPayload::StepStarted => "StepStarted", diff --git a/crates/merry-runtime/src/session/checkpoint_window/planning.rs b/crates/merry-runtime/src/session/checkpoint_window/planning.rs index 9266bf06..ec1c5e99 100644 --- a/crates/merry-runtime/src/session/checkpoint_window/planning.rs +++ b/crates/merry-runtime/src/session/checkpoint_window/planning.rs @@ -190,12 +190,13 @@ impl SessionState { &existing_open_archives, )?) .ok_or(CompactionError::BudgetOverflow)?; - let error = - if current_only_tokens >= window_budget.max_dynamic_body_tokens() { - CompactionError::UncompressibleCurrentInput - } else { - CompactionError::MinimumRawTurnCannotFit - }; + let error = if window_budget.estimate_body_tokens(current_only_tokens) + >= window_budget.max_dynamic_body_tokens() + { + CompactionError::UncompressibleCurrentInput + } else { + CompactionError::MinimumRawTurnCannotFit + }; last_failure = Some(error); } else if last_failure.is_none() { last_failure = Some(CompactionError::NoWindowFitsCompactionRequest); @@ -217,7 +218,9 @@ impl SessionState { .archive_only_fixed_dynamic_body_tokens() .checked_add(projected_turn_tokens(open_turns, &existing_open_archives)?) .ok_or(CompactionError::BudgetOverflow)?; - if current_only_tokens >= window_budget.max_dynamic_body_tokens() { + if window_budget.estimate_body_tokens(current_only_tokens) + >= window_budget.max_dynamic_body_tokens() + { return Err(CompactionError::UncompressibleCurrentInput.into()); } } @@ -295,7 +298,7 @@ fn plan_retained_window( base_tokens, raw_turns, archived_tool_call_ids, - window_budget.max_dynamic_body_tokens(), + window_budget, ) }; @@ -416,12 +419,12 @@ pub(super) fn retained_projection_fits( base_tokens: u64, raw_turns: &[ModelTurnHistory<'_>], archived_tool_call_ids: &BTreeSet, - max_dynamic_body_tokens: u64, + window_budget: CompactionWindowBudget, ) -> Result { - Ok(base_tokens + let base_tokens = base_tokens .checked_add(projected_turn_tokens(raw_turns, archived_tool_call_ids)?) - .ok_or(CompactionError::BudgetOverflow)? - < max_dynamic_body_tokens) + .ok_or(CompactionError::BudgetOverflow)?; + Ok(window_budget.estimate_body_tokens(base_tokens) < window_budget.max_dynamic_body_tokens()) } pub(super) fn compaction_window_plan( diff --git a/crates/merry-runtime/src/session/mod.rs b/crates/merry-runtime/src/session/mod.rs index d5e79231..5dc8ae91 100644 --- a/crates/merry-runtime/src/session/mod.rs +++ b/crates/merry-runtime/src/session/mod.rs @@ -66,6 +66,7 @@ pub(crate) struct SessionState { pending_tool_calls: Vec, resolved_tool_calls: BTreeSet, usage: Option, + request_token_calibration: Option, trajectory_snapshot: Option, external_tool_catalog: merry_core::SessionToolCatalog, } @@ -101,6 +102,7 @@ impl SessionState { pending_tool_calls: Vec::new(), resolved_tool_calls: BTreeSet::new(), usage: None, + request_token_calibration: None, trajectory_snapshot: None, external_tool_catalog, } diff --git a/crates/merry-runtime/src/session/persistence.rs b/crates/merry-runtime/src/session/persistence.rs index edd163ad..874a294b 100644 --- a/crates/merry-runtime/src/session/persistence.rs +++ b/crates/merry-runtime/src/session/persistence.rs @@ -62,6 +62,8 @@ struct StoredSessionDocument { transcript: PersistedTranscript, resolved_tool_calls: Vec, usage: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + request_token_calibration: Option, task_anchor: Option, registries: StoredRegistries, active_plan: Option, @@ -320,6 +322,7 @@ impl SessionState { transcript: view.transcript.persisted(), resolved_tool_calls: view.resolved_tool_calls.iter().cloned().collect(), usage: self.usage.clone(), + request_token_calibration: self.request_token_calibration.clone(), task_anchor: self.task_anchor.as_ref().map(|anchor| StoredTaskAnchor { objective: anchor.objective().to_owned(), }), @@ -416,6 +419,7 @@ impl SessionState { .into_iter() .collect::>(), usage: document.usage, + request_token_calibration: document.request_token_calibration, trajectory_snapshot: document.trajectory_snapshot, external_tool_catalog: document.external_tool_catalog, }; diff --git a/crates/merry-runtime/src/session/usage.rs b/crates/merry-runtime/src/session/usage.rs index adba6500..e2f4dfb0 100644 --- a/crates/merry-runtime/src/session/usage.rs +++ b/crates/merry-runtime/src/session/usage.rs @@ -1,11 +1,38 @@ use super::SessionState; use crate::ledger::LedgerFactKind; +use crate::token_estimate::{RequestTokenCalibration, RequestTokenObservation, TokenEstimateScale}; use merry_core::{ CompactionUsageWindow, ErrorInfo, ModelUsage, RuntimeJournalEvent, RuntimeJournalPayload, SessionUsage, UsageContextWindow, }; impl SessionState { + pub(crate) fn token_estimate_scale( + &self, + provider: &merry_core::ProviderName, + request: &merry_llm::ModelRequest, + ) -> TokenEstimateScale { + self.request_token_calibration + .as_ref() + .map_or_else(TokenEstimateScale::default, |calibration| { + calibration.scale_for(provider, request) + }) + } + + pub(crate) fn calibrate_request_tokens( + &mut self, + observation: RequestTokenObservation, + usage: ModelUsage, + ) { + if let Some(calibration) = RequestTokenCalibration::observe( + self.request_token_calibration.as_ref(), + observation, + usage.input_tokens(), + ) { + self.request_token_calibration = Some(calibration); + } + } + pub(crate) fn usage(&self) -> Option { self.usage.clone() } diff --git a/crates/merry-runtime/src/token_estimate.rs b/crates/merry-runtime/src/token_estimate.rs index cdde6ad4..6332f2f4 100644 --- a/crates/merry-runtime/src/token_estimate.rs +++ b/crates/merry-runtime/src/token_estimate.rs @@ -2,6 +2,18 @@ use merry_llm::{ModelContent, ModelInputItem, ModelRequest, ModelResponseFormat}; +mod calibration; + +pub(crate) use calibration::{ + RequestTokenCalibration, RequestTokenObservation, TokenEstimateScale, +}; + +/// Estimates the complete input, including tools and response schemas, before calibration. +pub(crate) fn estimate_request_input_tokens(request: &ModelRequest) -> u64 { + estimate_model_input_tokens(request.input()) + .saturating_add(estimate_request_contract_tokens(request)) +} + /// Estimates provider-visible tools and response schemas in addition to messages. pub(crate) fn estimate_request_contract_tokens(request: &ModelRequest) -> u64 { let tools = request @@ -22,13 +34,13 @@ pub(crate) fn estimate_request_contract_tokens(request: &ModelRequest) -> u64 { tools.saturating_add(format) } -/// Bytes per token used by every text estimate in the runtime. +/// Bytes per token used by the deterministic base text estimate. /// -/// Budgets, window fitting, and planning all compare against this one ratio, so -/// a change here moves all of them together. It is deliberately the optimistic -/// axis of the estimate, distinct from compaction's accepted-output byte -/// ceiling ([`crate::compaction`]), which adds slack so a checkpoint that fits -/// the token budget is not rejected on byte count. +/// Primary request budgets and destination-history planning apply session-owned +/// usage calibration on top. Independent compactor requests do not inherit that +/// model's multiplier. This remains distinct from compaction's accepted-output +/// byte ceiling ([`crate::compaction`]), which adds slack so a checkpoint that +/// fits the token budget is not rejected on byte count. pub(crate) const BYTES_PER_TOKEN: u64 = 4; pub(crate) fn estimate_model_input_tokens(input: &[ModelInputItem]) -> u64 { @@ -36,9 +48,11 @@ pub(crate) fn estimate_model_input_tokens(input: &[ModelInputItem]) -> u64 { } pub(crate) fn estimate_text_tokens(text: &str) -> u64 { - u64::try_from(text.len()) - .expect("usize should fit in u64 on supported targets") - .div_ceil(BYTES_PER_TOKEN) + estimate_utf8_tokens(u64::try_from(text.len()).unwrap_or(u64::MAX)) +} + +pub(crate) const fn estimate_utf8_tokens(bytes: u64) -> u64 { + bytes.div_ceil(BYTES_PER_TOKEN) } fn estimate_model_input_item_tokens(item: &ModelInputItem) -> u64 { diff --git a/crates/merry-runtime/src/token_estimate/calibration.rs b/crates/merry-runtime/src/token_estimate/calibration.rs new file mode 100644 index 00000000..9c0acd73 --- /dev/null +++ b/crates/merry-runtime/src/token_estimate/calibration.rs @@ -0,0 +1,151 @@ +//! Session-local feedback from complete request estimates and matching provider usage. + +use super::estimate_request_input_tokens; +use merry_core::ProviderName; +use merry_llm::{ModelInputItem, ModelName, ModelRequest, RequestContentHash}; +use serde::{Deserialize, Serialize}; +use std::num::NonZeroU64; + +const SCALE_PRECISION: u64 = 1_000_000; + +/// Fixed-point input multiplier; estimates round upward and saturate on overflow. +/// An overflowing ratio is stored as the maximum value and fails closed on estimation. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(transparent)] +pub(crate) struct TokenEstimateScale(NonZeroU64); + +impl Default for TokenEstimateScale { + fn default() -> Self { + Self(NonZeroU64::MIN.saturating_add(SCALE_PRECISION - 1)) + } +} + +impl TokenEstimateScale { + /// Converts a base estimate to the calibrated input-token domain. + pub(crate) fn estimate(self, base_tokens: u64) -> u64 { + if base_tokens != 0 && self.0 == NonZeroU64::MAX { + return u64::MAX; + } + let tokens = (u128::from(base_tokens) * u128::from(self.0.get())) + .div_ceil(u128::from(SCALE_PRECISION)); + u64::try_from(tokens).unwrap_or(u64::MAX) + } + + fn from_usage(base_tokens: u64, actual_tokens: u64) -> Option { + if base_tokens == 0 || actual_tokens == 0 { + return None; + } + let ratio = (u128::from(actual_tokens) * u128::from(SCALE_PRECISION)) + .div_ceil(u128::from(base_tokens)); + NonZeroU64::new(u64::try_from(ratio).unwrap_or(u64::MAX)).map(Self) + } + + /// Corrects underestimation immediately; releases excess headroom gradually. + fn updated(self, sample: Self) -> Self { + if sample.0 >= self.0 { + return sample; + } + let smoothed = self.0.get() - (self.0.get() - sample.0.get()) / 4; + NonZeroU64::new(smoothed).map_or(sample, Self) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +struct RequestTokenProfile { + provider: ProviderName, + model: ModelName, + stable_prefix: RequestContentHash, + has_images: bool, +} + +impl RequestTokenProfile { + fn new(provider: &ProviderName, request: &ModelRequest) -> Self { + Self { + provider: provider.clone(), + model: request.model().clone(), + stable_prefix: request.stable_prefix_hash().clone(), + has_images: request_has_images(request), + } + } + + fn matches(&self, provider: &ProviderName, request: &ModelRequest) -> bool { + self.provider == *provider + && self.model == *request.model() + && self.stable_prefix == *request.stable_prefix_hash() + && self.has_images == request_has_images(request) + } +} + +fn request_has_images(request: &ModelRequest) -> bool { + request.input().iter().any(|item| { + matches!(item, ModelInputItem::Message(message) if message.content().images().next().is_some()) + }) +} + +/// Captured after compaction and before sending, so usage cannot match a stale estimate. +pub(crate) struct RequestTokenObservation { + profile: RequestTokenProfile, + base_input_tokens: u64, +} + +impl RequestTokenObservation { + pub(crate) fn new(provider: &ProviderName, request: &ModelRequest) -> Self { + Self { + profile: RequestTokenProfile::new(provider, request), + base_input_tokens: estimate_request_input_tokens(request), + } + } +} + +/// Bounded to one active request profile; unrelated models and contracts start uncalibrated. +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub(crate) struct RequestTokenCalibration { + profile: RequestTokenProfile, + scale: TokenEstimateScale, +} + +impl RequestTokenCalibration { + pub(crate) fn scale_for( + &self, + provider: &ProviderName, + request: &ModelRequest, + ) -> TokenEstimateScale { + if self.profile.matches(provider, request) { + self.scale + } else { + TokenEstimateScale::default() + } + } + + /// Returns no update for zero measurements; actual input includes cached tokens. + pub(crate) fn observe( + previous: Option<&Self>, + observation: RequestTokenObservation, + actual_input_tokens: u64, + ) -> Option { + let sample = + TokenEstimateScale::from_usage(observation.base_input_tokens, actual_input_tokens)?; + let previous_scale = previous + .filter(|previous| previous.profile == observation.profile) + .map_or_else(TokenEstimateScale::default, |previous| previous.scale); + let scale = previous_scale.updated(sample); + tracing::debug!( + event = "runtime.context.estimate_calibrated", + provider = observation.profile.provider.as_str(), + model = observation.profile.model.as_str(), + base_input_tokens = observation.base_input_tokens, + actual_input_tokens, + scale_parts_per_million = scale.0.get(), + "updated primary request token estimate calibration" + ); + Some(Self { + profile: observation.profile, + scale, + }) + } +} + +#[cfg(test)] +mod tests; diff --git a/crates/merry-runtime/src/token_estimate/calibration/tests.rs b/crates/merry-runtime/src/token_estimate/calibration/tests.rs new file mode 100644 index 00000000..4a060f9f --- /dev/null +++ b/crates/merry-runtime/src/token_estimate/calibration/tests.rs @@ -0,0 +1,235 @@ +use super::{RequestTokenCalibration, RequestTokenObservation, TokenEstimateScale}; +use crate::token_estimate::{estimate_model_input_tokens, estimate_request_input_tokens}; +use merry_core::{ProviderName, ToolInputSchema, ToolName, ToolSpec}; +use merry_llm::{ + GenerationConfig, ModelContent, ModelMessage, ModelMessageRole, ModelName, ModelRequest, +}; +use serde_json::json; + +fn provider(name: &str) -> ProviderName { + ProviderName::new(name).expect("provider") +} + +fn request(model: &str, instructions: &str, text: &str, tool_description: &str) -> ModelRequest { + let tool = ToolSpec::new( + ToolName::new("lookup").expect("tool name"), + tool_description, + ToolInputSchema::new( + schemars::Schema::try_from(json!({"type": "object", "properties": {}})) + .expect("schema"), + ) + .expect("input schema"), + ) + .expect("tool"); + ModelRequest::new_with_continuations_and_stable_prefix( + ModelName::new(model).expect("model"), + vec![ + ModelMessage::new( + ModelMessageRole::System, + ModelContent::text(instructions).expect("text"), + ) + .expect("instructions"), + ModelMessage::new( + ModelMessageRole::User, + ModelContent::text(text).expect("text"), + ) + .expect("input"), + ], + vec![tool], + Vec::new(), + GenerationConfig::default(), + 1, + ) + .expect("request") +} + +#[test] +fn feedback_matches_the_complete_request_not_just_the_dynamic_body() { + let provider = provider("provider"); + let request = request( + "model", + &"rules ".repeat(1_000), + "hello", + &"tool ".repeat(100), + ); + let base_tokens = estimate_request_input_tokens(&request); + assert!(base_tokens > estimate_model_input_tokens(request.dynamic_input()) * 100); + let calibration = RequestTokenCalibration::observe( + None, + RequestTokenObservation::new(&provider, &request), + base_tokens * 2, + ) + .expect("measurement"); + assert_eq!( + calibration.scale_for(&provider, &request).estimate(10_000), + 20_000 + ); +} + +#[test] +fn rising_input_cost_is_corrected_immediately_and_falling_cost_is_smoothed() { + let provider = provider("provider"); + let request = request("model", "rules", "hello", "lookup"); + let base_tokens = estimate_request_input_tokens(&request); + let mut calibration = None; + for multiplier in [2, 3, 4] { + calibration = RequestTokenCalibration::observe( + calibration.as_ref(), + RequestTokenObservation::new(&provider, &request), + base_tokens * multiplier, + ); + assert_eq!( + calibration + .as_ref() + .expect("calibration") + .scale_for(&provider, &request) + .estimate(base_tokens), + base_tokens * multiplier, + ); + } + let updated = RequestTokenCalibration::observe( + calibration.as_ref(), + RequestTokenObservation::new(&provider, &request), + base_tokens * 2, + ) + .expect("measurement"); + let estimate = updated.scale_for(&provider, &request).estimate(base_tokens); + assert!(estimate > base_tokens * 2 && estimate < base_tokens * 4); +} + +#[test] +fn consistent_overestimation_converges_without_an_abrupt_drop() { + let provider = provider("provider"); + let request = request("model", "rules", &"abcd".repeat(1_000), "lookup"); + let base_tokens = estimate_request_input_tokens(&request); + let actual_tokens = base_tokens / 2; + let mut calibration = None; + let mut previous = base_tokens; + for _ in 0..32 { + calibration = RequestTokenCalibration::observe( + calibration.as_ref(), + RequestTokenObservation::new(&provider, &request), + actual_tokens, + ); + let estimated = calibration + .as_ref() + .expect("calibration") + .scale_for(&provider, &request) + .estimate(base_tokens); + assert!((actual_tokens..=previous).contains(&estimated)); + previous = estimated; + } + assert!(previous <= actual_tokens + 1); +} + +#[test] +fn unrelated_provider_model_prefix_and_tool_contracts_do_not_share_feedback() { + let source_provider = provider("primary"); + let source = request("model", "rules", "first", "lookup"); + let calibration = RequestTokenCalibration::observe( + None, + RequestTokenObservation::new(&source_provider, &source), + estimate_request_input_tokens(&source) * 2, + ) + .expect("measurement"); + for (provider, candidate) in [ + ( + provider("other"), + request("model", "rules", "next", "lookup"), + ), + ( + provider("primary"), + request("other-model", "rules", "next", "lookup"), + ), + ( + provider("primary"), + request("model", "new rules", "next", "lookup"), + ), + ( + provider("primary"), + request("model", "rules", "next", "new tool"), + ), + ] { + assert_eq!( + calibration.scale_for(&provider, &candidate).estimate(100), + 100 + ); + } + let next = request("model", "rules", "different dynamic input", "lookup"); + assert_eq!( + calibration.scale_for(&source_provider, &next).estimate(100), + 200 + ); +} + +#[test] +fn image_requests_do_not_reuse_text_only_feedback() { + let provider = provider("provider"); + let request = request("model", "rules", "hello", "lookup"); + let calibration = RequestTokenCalibration::observe( + None, + RequestTokenObservation::new(&provider, &request), + estimate_request_input_tokens(&request) * 2, + ) + .expect("measurement"); + let image = merry_llm::ModelImage::png( + "[Image #1]", + std::sync::Arc::<[u8]>::from([137, 80, 78, 71, 13, 10, 26, 10]), + 100, + 100, + ) + .expect("image"); + let mut input = request.input().to_vec(); + input.push(merry_llm::ModelInputItem::Message( + ModelMessage::new( + ModelMessageRole::User, + ModelContent::user_with_images("inspect [Image #1]", vec![image]) + .expect("image content"), + ) + .expect("message"), + )); + let image_request = ModelRequest::new_with_input_and_stable_prefix( + request.model().clone(), + input, + request.tools().to_vec(), + GenerationConfig::default(), + request.stable_prefix_item_count(), + ) + .expect("image request"); + assert_eq!( + calibration + .scale_for(&provider, &image_request) + .estimate(100), + 100 + ); +} + +#[test] +fn zero_measurements_are_ignored_and_large_counts_saturate_safely() { + assert!(TokenEstimateScale::from_usage(0, 100).is_none()); + assert!(TokenEstimateScale::from_usage(100, 0).is_none()); + let scale = TokenEstimateScale::from_usage(3, 4).expect("scale"); + assert_eq!(scale.estimate(0), 0); + assert_eq!(scale.estimate(1), 2); + let scale = TokenEstimateScale::from_usage(1, u64::MAX).expect("scale"); + assert_eq!(scale.estimate(1), u64::MAX); + assert_eq!(scale.estimate(u64::MAX), u64::MAX); +} + +#[test] +fn persisted_feedback_round_trips_and_rejects_a_zero_scale() { + let provider = provider("provider"); + let request = request("model", "rules", "hello", "lookup"); + let calibration = RequestTokenCalibration::observe( + None, + RequestTokenObservation::new(&provider, &request), + estimate_request_input_tokens(&request) * 2, + ) + .expect("measurement"); + let mut encoded = serde_json::to_value(&calibration).expect("serialize calibration"); + let decoded: RequestTokenCalibration = + serde_json::from_value(encoded.clone()).expect("restore calibration"); + assert_eq!(decoded.scale_for(&provider, &request).estimate(100), 200); + encoded["scale"] = json!(0); + assert!(serde_json::from_value::(encoded).is_err()); +} diff --git a/crates/merry-runtime/tests/agent_loop.rs b/crates/merry-runtime/tests/agent_loop.rs index 4983cb55..0a8c80e0 100644 --- a/crates/merry-runtime/tests/agent_loop.rs +++ b/crates/merry-runtime/tests/agent_loop.rs @@ -18,6 +18,8 @@ mod mutation_admission; mod stream_lifecycle; #[path = "agent_loop/support/mod.rs"] mod support; +#[path = "agent_loop/telemetry_retention.rs"] +mod telemetry_retention; #[path = "agent_loop/tool_scheduling.rs"] mod tool_scheduling; #[path = "agent_loop/tracing.rs"] diff --git a/crates/merry-runtime/tests/agent_loop/support/events.rs b/crates/merry-runtime/tests/agent_loop/support/events.rs index 821cd2bc..6b4d4483 100644 --- a/crates/merry-runtime/tests/agent_loop/support/events.rs +++ b/crates/merry-runtime/tests/agent_loop/support/events.rs @@ -39,6 +39,12 @@ pub(crate) fn assert_sanitized_policy_denial_json(value: &Value, tool_name: &str pub(crate) fn event_kind_names(events: &[RuntimeJournalEvent]) -> Vec<&'static str> { events .iter() + .filter(|event| { + !matches!( + event.payload, + RuntimeJournalPayload::ModelOutputRateUpdated { .. } + ) + }) .map(|event| match event.payload { RuntimeJournalPayload::SessionStarted => "SessionStarted", RuntimeJournalPayload::StepStarted => "StepStarted", @@ -66,6 +72,7 @@ pub(crate) fn event_kind_names(events: &[RuntimeJournalEvent]) -> Vec<&'static s pub(crate) fn public_event_kind_names(events: &[RuntimeEvent]) -> Vec<&'static str> { events .iter() + .filter(|event| !matches!(event, RuntimeEvent::ModelOutputRateUpdated { .. })) .map(|event| match event { RuntimeEvent::SessionStarted { .. } => "SessionStarted", RuntimeEvent::StepStarted { .. } => "StepStarted", diff --git a/crates/merry-runtime/tests/agent_loop/telemetry_retention.rs b/crates/merry-runtime/tests/agent_loop/telemetry_retention.rs new file mode 100644 index 00000000..9077f43c --- /dev/null +++ b/crates/merry-runtime/tests/agent_loop/telemetry_retention.rs @@ -0,0 +1,87 @@ +use crate::support::{ + models::{ScriptedModelProvider, completed_text_event}, + runtime::{run_default_loop, runtime_with_provider}, +}; +use merry_core::RuntimeEvent; +use merry_llm::{ModelError, ModelEvent, ModelOutputProgress}; +use merry_runtime::{AgentLoopConfig, AgentRunMessage, StepContext, StepInput}; +use std::time::Duration; + +fn long_output(fragments: u64) -> Vec> { + (0..fragments) + .flat_map(|index| { + [ + Ok(ModelEvent::OutputProgress { + progress: Some(ModelOutputProgress::new( + index * 4, + Duration::from_millis(index), + )), + }), + Ok(ModelEvent::OutputTextDelta { + delta: "x".to_owned(), + }), + ] + }) + .chain([Ok(completed_text_event("done"))]) + .collect() +} + +#[tokio::test] +async fn retained_run_evidence_does_not_grow_with_transient_output_fragments() { + let mut retained_count = None; + for fragments in [1, 10_000] { + let runtime = runtime_with_provider( + "bounded-retained-telemetry", + ScriptedModelProvider::new(vec![long_output(fragments)]), + ); + let result = run_default_loop(&runtime, "finish").await; + assert_eq!(result.final_output(), Some("done")); + assert!( + result + .events() + .iter() + .all(|event| !event.payload.is_transient()) + ); + if let Some(count) = retained_count { + assert_eq!(result.events().len(), count); + } + retained_count = Some(result.events().len()); + } +} + +#[tokio::test] +async fn live_output_is_forwarded_without_entering_the_stream_result() { + let runtime = runtime_with_provider( + "bounded-stream-telemetry", + ScriptedModelProvider::new(vec![long_output(10_000)]), + ); + let mut run = runtime + .run_agent_loop_stream( + StepInput::user_text("finish").unwrap(), + StepContext::default(), + AgentLoopConfig::default(), + ) + .unwrap(); + let mut deltas = 0; + let mut samples = 0; + while let Some(message) = run.next_message().await.unwrap() { + match message { + AgentRunMessage::Event(RuntimeEvent::AssistantMessageDelta { .. }) => deltas += 1, + AgentRunMessage::Event(RuntimeEvent::ModelOutputRateUpdated { + rate: Some(_), .. + }) => samples += 1, + _ => {} + } + } + let result = run.result().await.unwrap(); + assert_eq!(deltas, 10_000); + assert!(samples > 0); + assert_eq!(result.final_output(), Some("done")); + assert!( + result + .events() + .iter() + .all(|event| !event.payload.is_transient()) + ); + assert!(result.events().len() < 20); +} diff --git a/crates/merry-runtime/tests/agent_loop/tool_scheduling.rs b/crates/merry-runtime/tests/agent_loop/tool_scheduling.rs index 7ea69c0a..ac699196 100644 --- a/crates/merry-runtime/tests/agent_loop/tool_scheduling.rs +++ b/crates/merry-runtime/tests/agent_loop/tool_scheduling.rs @@ -139,7 +139,7 @@ async fn agent_loop_executes_one_tool_and_continues_to_final_completion() { .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2, 3, 4, 5, 6, 7] + vec![0, 1, 3, 4, 5, 6, 8, 9] ); assert!(runtime.pending_tool_calls().await.is_empty()); assert_eq!(executor.calls().len(), 1); diff --git a/crates/merry-runtime/tests/cancellation.rs b/crates/merry-runtime/tests/cancellation.rs index fb3b8a46..e503289e 100644 --- a/crates/merry-runtime/tests/cancellation.rs +++ b/crates/merry-runtime/tests/cancellation.rs @@ -82,6 +82,12 @@ async fn collect_pending_step(runtime: &Runtime, text: &str) -> Vec Vec<&'static str> { events .iter() + .filter(|event| { + !matches!( + event.payload, + RuntimeJournalPayload::ModelOutputRateUpdated { .. } + ) + }) .map(|event| match event.payload { RuntimeJournalPayload::SessionStarted => "SessionStarted", RuntimeJournalPayload::StepStarted => "StepStarted", @@ -104,7 +110,7 @@ fn event_kind_names(events: &[RuntimeJournalEvent]) -> Vec<&'static str> { async fn assert_missing_tool_result_artifact(runtime: &Runtime) { let evidence_err = runtime .evidence_ref( - &artifact_id("tool-result-3"), + &artifact_id("tool-result-4"), EvidenceLocator::whole_artifact(), ) .await @@ -113,7 +119,7 @@ async fn assert_missing_tool_result_artifact(runtime: &Runtime) { evidence_err, RuntimeError::Artifact { source: ArtifactError::MissingArtifact { id } - } if id == artifact_id("tool-result-3") + } if id == artifact_id("tool-result-4") )); } @@ -308,7 +314,7 @@ async fn pre_cancelled_tool_execution_keeps_pending_and_releases_active_permit() .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2] + vec![0, 1, 2, 3] ); let projection_before_cancel = runtime.ledger_projection().await; assert_eq!( @@ -325,7 +331,7 @@ async fn pre_cancelled_tool_execution_keeps_pending_and_releases_active_permit() kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, @@ -409,7 +415,7 @@ async fn cancelling_during_tool_execution_keeps_pending_and_releases_active_perm .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2] + vec![0, 1, 2, 3] ); let projection_before_cancel = runtime.ledger_projection().await; assert_eq!( @@ -426,7 +432,7 @@ async fn cancelling_during_tool_execution_keeps_pending_and_releases_active_perm kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, @@ -483,7 +489,7 @@ async fn cancelling_after_successful_tool_execution_keeps_pending_and_releases_a .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2] + vec![0, 1, 2, 3] ); let pending_before_cancel = runtime.pending_tool_calls().await; assert_eq!(pending_before_cancel.len(), 1); @@ -502,7 +508,7 @@ async fn cancelling_after_successful_tool_execution_keeps_pending_and_releases_a kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, diff --git a/crates/merry-runtime/tests/provider_boundary/compaction_semantics.rs b/crates/merry-runtime/tests/provider_boundary/compaction_semantics.rs index 98c65234..5fa37878 100644 --- a/crates/merry-runtime/tests/provider_boundary/compaction_semantics.rs +++ b/crates/merry-runtime/tests/provider_boundary/compaction_semantics.rs @@ -393,7 +393,7 @@ async fn collect_text_output(stream: ModelEventStream) -> Result let mut saw_delta = false; while let Some(item) = stream.next().await { match item { - Ok(ModelEvent::Started) => {} + Ok(ModelEvent::Started | ModelEvent::OutputProgress { .. }) => {} Ok(ModelEvent::OutputTextDelta { delta }) => { if !delta.is_empty() { saw_delta = true; diff --git a/crates/merry-runtime/tests/provider_boundary/diagnostics.rs b/crates/merry-runtime/tests/provider_boundary/diagnostics.rs index 3f7cf434..a4206780 100644 --- a/crates/merry-runtime/tests/provider_boundary/diagnostics.rs +++ b/crates/merry-runtime/tests/provider_boundary/diagnostics.rs @@ -107,7 +107,7 @@ async fn t6_failure_evidence_links_journal_artifact_ledger_and_trajectory() { .iter() .map(|event| event.sequence) .collect::>(), - vec![3, 4] + vec![4, 5] ); assert!(matches!( &events[0].payload, @@ -148,17 +148,17 @@ async fn t6_failure_evidence_links_journal_artifact_ledger_and_trajectory() { kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, order: 3, kind: LedgerFactKind::ArtifactRecorded, }, LedgerProjection::Lifecycle { - sequence: 4, + sequence: 5, order: 4, kind: LedgerFactKind::ToolCallResolved, }, diff --git a/crates/merry-runtime/tests/provider_boundary/step_lifecycle.rs b/crates/merry-runtime/tests/provider_boundary/step_lifecycle.rs index f98b46a9..c5a82a88 100644 --- a/crates/merry-runtime/tests/provider_boundary/step_lifecycle.rs +++ b/crates/merry-runtime/tests/provider_boundary/step_lifecycle.rs @@ -54,10 +54,10 @@ async fn runtime_step_with_provider_compiles_user_text_request_and_records_assis .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2, 3, 4] + vec![0, 1, 2, 3, 4, 5] ); let artifact = assistant_output_artifact(&events); - assert_eq!(artifact.id().as_str(), "assistant-output-3"); + assert_eq!(artifact.id().as_str(), "assistant-output-4"); assert_eq!(artifact.kind(), &ArtifactKind::Text); let evidence = runtime .evidence_ref(artifact.id(), EvidenceLocator::whole_artifact()) @@ -80,12 +80,12 @@ async fn runtime_step_with_provider_compiles_user_text_request_and_records_assis kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, order: 2, kind: LedgerFactKind::ArtifactRecorded, }, LedgerProjection::Lifecycle { - sequence: 4, + sequence: 5, order: 3, kind: LedgerFactKind::StepCompleted, }, @@ -236,7 +236,7 @@ async fn reserved_assistant_output_external_recording_does_not_block_runtime_own ] ); let generated = assistant_output_artifact(&events); - assert_eq!(generated.id().as_str(), "assistant-output-2"); + assert_eq!(generated.id().as_str(), "assistant-output-3"); let evidence = runtime .evidence_ref(generated.id(), EvidenceLocator::whole_artifact()) .await @@ -407,14 +407,14 @@ async fn second_provider_step_continues_sequences_and_replays_transcript() { .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2, 3] + vec![0, 1, 2, 3, 4] ); assert_eq!( second_events .iter() .map(|event| event.sequence) .collect::>(), - vec![4, 5, 6] + vec![5, 6, 7, 8] ); assert_eq!( event_kind_names(&second_events), @@ -422,11 +422,11 @@ async fn second_provider_step_continues_sequences_and_replays_transcript() { ); assert_eq!( assistant_output_artifact(&first_events).id().as_str(), - "assistant-output-2" + "assistant-output-3" ); assert_eq!( assistant_output_artifact(&second_events).id().as_str(), - "assistant-output-5" + "assistant-output-7" ); let requests = provider.recorded_requests(); diff --git a/crates/merry-runtime/tests/provider_boundary/stream_failures.rs b/crates/merry-runtime/tests/provider_boundary/stream_failures.rs index 6c5ddb5c..c4058f49 100644 --- a/crates/merry-runtime/tests/provider_boundary/stream_failures.rs +++ b/crates/merry-runtime/tests/provider_boundary/stream_failures.rs @@ -35,7 +35,7 @@ async fn provider_stream_error_emits_failed_without_step_completed() { let projection = runtime.ledger_projection().await; assert!(projection.entries().contains(&LedgerProjection::Lifecycle { sequence: failed_sequence, - order: failed_sequence, + order: 2, kind: LedgerFactKind::Failed })); } diff --git a/crates/merry-runtime/tests/provider_boundary/support/events.rs b/crates/merry-runtime/tests/provider_boundary/support/events.rs index 2ac13e77..df6db532 100644 --- a/crates/merry-runtime/tests/provider_boundary/support/events.rs +++ b/crates/merry-runtime/tests/provider_boundary/support/events.rs @@ -6,6 +6,12 @@ use merry_core::{ pub(crate) fn event_kind_names(events: &[RuntimeJournalEvent]) -> Vec<&'static str> { events .iter() + .filter(|event| { + !matches!( + event.payload, + RuntimeJournalPayload::ModelOutputRateUpdated { .. } + ) + }) .map(|event| match event.payload { RuntimeJournalPayload::SessionStarted => "SessionStarted", RuntimeJournalPayload::StepStarted => "StepStarted", diff --git a/crates/merry-runtime/tests/provider_boundary/tool_admission.rs b/crates/merry-runtime/tests/provider_boundary/tool_admission.rs index 0fb1a7df..2c1ed7ec 100644 --- a/crates/merry-runtime/tests/provider_boundary/tool_admission.rs +++ b/crates/merry-runtime/tests/provider_boundary/tool_admission.rs @@ -44,7 +44,7 @@ async fn unregistered_pending_tool_name_resolves_failed_with_tool_not_registered .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2] + vec![0, 1, 2, 3] ); let pending = pending_tool_call(&pending_events).clone(); @@ -62,7 +62,7 @@ async fn unregistered_pending_tool_name_resolves_failed_with_tool_not_registered .iter() .map(|event| event.sequence) .collect::>(), - vec![3, 4] + vec![4, 5] ); let result = resolved_tool_result(&execution_events); assert!(matches!( @@ -81,7 +81,7 @@ async fn unregistered_pending_tool_name_resolves_failed_with_tool_not_registered .code(), "tool_not_registered" ); - assert_eq!(result.artifact().id().as_str(), "tool-result-3"); + assert_eq!(result.artifact().id().as_str(), "tool-result-4"); assert_eq!(result.artifact().kind(), &ArtifactKind::Json); assert_eq!(result.call_id(), pending.id()); assert!(failed_code(&execution_events).is_none()); @@ -106,17 +106,17 @@ async fn unregistered_pending_tool_name_resolves_failed_with_tool_not_registered kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, order: 3, kind: LedgerFactKind::ArtifactRecorded, }, LedgerProjection::Lifecycle { - sequence: 4, + sequence: 5, order: 4, kind: LedgerFactKind::ToolCallResolved, }, @@ -133,7 +133,7 @@ async fn unregistered_pending_tool_name_resolves_failed_with_tool_not_registered .iter() .map(|event| event.sequence) .collect::>(), - vec![5, 6, 7] + vec![6, 7, 8, 9] ); assert_eq!(provider.recorded_requests()[1].continuations().len(), 1); } diff --git a/crates/merry-runtime/tests/provider_boundary/tool_evidence.rs b/crates/merry-runtime/tests/provider_boundary/tool_evidence.rs index dd718d02..9477d75c 100644 --- a/crates/merry-runtime/tests/provider_boundary/tool_evidence.rs +++ b/crates/merry-runtime/tests/provider_boundary/tool_evidence.rs @@ -185,7 +185,7 @@ async fn execute_registered_tool_success_records_artifact_resolves_and_compiles_ .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2] + vec![0, 1, 2, 3] ); let pending = pending_tool_call(&pending_events).clone(); let reserved_artifact = ArtifactRef::new(artifact_id("tool-result-4"), ArtifactKind::Text); @@ -231,7 +231,7 @@ async fn execute_registered_tool_success_records_artifact_resolves_and_compiles_ .iter() .map(|event| event.sequence) .collect::>(), - vec![3, 4] + vec![4, 5] ); let result = resolved_tool_result(&execution_events); assert!(matches!( @@ -243,7 +243,7 @@ async fn execute_registered_tool_success_records_artifact_resolves_and_compiles_ RuntimeJournalPayload::ToolCallResolved { result: resolved } if resolved == result )); assert_eq!(result.status(), ToolCallResultStatus::Succeeded); - assert_eq!(result.artifact().id().as_str(), "tool-result-3"); + assert_eq!(result.artifact().id().as_str(), "tool-result-4"); assert_eq!(result.artifact().kind(), &ArtifactKind::Text); assert_eq!(result.call_id(), pending.id()); let evidence = runtime @@ -267,17 +267,17 @@ async fn execute_registered_tool_success_records_artifact_resolves_and_compiles_ kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, order: 3, kind: LedgerFactKind::ArtifactRecorded, }, LedgerProjection::Lifecycle { - sequence: 4, + sequence: 5, order: 4, kind: LedgerFactKind::ToolCallResolved, }, @@ -294,7 +294,7 @@ async fn execute_registered_tool_success_records_artifact_resolves_and_compiles_ .iter() .map(|event| event.sequence) .collect::>(), - vec![5, 6, 7] + vec![6, 7, 8, 9] ); let requests = provider.recorded_requests(); @@ -370,7 +370,7 @@ async fn reading_catalog_skill_file_emits_skill_used_event() { .iter() .map(|event| event.sequence) .collect::>(), - vec![3, 4, 5] + vec![4, 5, 6] ); let result = resolved_tool_result(&execution_events); assert!(matches!( @@ -406,7 +406,7 @@ async fn execute_tool_domain_failure_resolves_failed_without_runtime_failed() { .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2] + vec![0, 1, 2, 3] ); let pending = pending_tool_call(&pending_events).clone(); @@ -424,7 +424,7 @@ async fn execute_tool_domain_failure_resolves_failed_without_runtime_failed() { .iter() .map(|event| event.sequence) .collect::>(), - vec![3, 4] + vec![4, 5] ); assert!( execution_events @@ -442,7 +442,7 @@ async fn execute_tool_domain_failure_resolves_failed_without_runtime_failed() { RuntimeJournalPayload::ToolCallResolved { result: resolved } if resolved == result )); assert_eq!(result.status(), ToolCallResultStatus::Failed); - assert_eq!(result.artifact().id().as_str(), "tool-result-3"); + assert_eq!(result.artifact().id().as_str(), "tool-result-4"); assert_eq!(result.artifact().kind(), &ArtifactKind::Json); assert_eq!(result.call_id(), pending.id()); assert_eq!( @@ -473,17 +473,17 @@ async fn execute_tool_domain_failure_resolves_failed_without_runtime_failed() { kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, order: 3, kind: LedgerFactKind::ArtifactRecorded, }, LedgerProjection::Lifecycle { - sequence: 4, + sequence: 5, order: 4, kind: LedgerFactKind::ToolCallResolved, }, diff --git a/crates/merry-runtime/tests/provider_boundary/tool_failures.rs b/crates/merry-runtime/tests/provider_boundary/tool_failures.rs index 3b99419c..ab2a8a88 100644 --- a/crates/merry-runtime/tests/provider_boundary/tool_failures.rs +++ b/crates/merry-runtime/tests/provider_boundary/tool_failures.rs @@ -132,7 +132,7 @@ async fn executor_infrastructure_error_keeps_pending_without_artifact_or_result( .iter() .map(|event| event.sequence) .collect::>(), - vec![0, 1, 2] + vec![0, 1, 2, 3] ); let pending = pending_tool_call(&pending_events).clone(); let before = runtime.ledger_projection().await; @@ -150,7 +150,7 @@ async fn executor_infrastructure_error_keeps_pending_without_artifact_or_result( kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, @@ -172,7 +172,7 @@ async fn executor_infrastructure_error_keeps_pending_without_artifact_or_result( assert_eq!(runtime.pending_tool_calls().await, vec![pending]); let evidence_err = runtime .evidence_ref( - &artifact_id("tool-result-3"), + &artifact_id("tool-result-4"), EvidenceLocator::whole_artifact(), ) .await @@ -181,7 +181,7 @@ async fn executor_infrastructure_error_keeps_pending_without_artifact_or_result( evidence_err, merry_runtime::RuntimeError::Artifact { source: ArtifactError::MissingArtifact { id } - } if id == artifact_id("tool-result-3") + } if id == artifact_id("tool-result-4") )); } @@ -213,13 +213,13 @@ async fn execute_tool_blank_text_outcome_keeps_pending_without_artifact_or_resul merry_runtime::RuntimeError::UnsupportedToolResultContent { artifact_id, content_kind: ArtifactContentKind::Text - } if artifact_id.as_str() == "tool-result-3" + } if artifact_id.as_str() == "tool-result-4" )); assert_eq!(before, after); assert_eq!(runtime.pending_tool_calls().await, vec![pending]); let evidence_err = runtime .evidence_ref( - &artifact_id("tool-result-3"), + &artifact_id("tool-result-4"), EvidenceLocator::whole_artifact(), ) .await @@ -228,7 +228,7 @@ async fn execute_tool_blank_text_outcome_keeps_pending_without_artifact_or_resul evidence_err, merry_runtime::RuntimeError::Artifact { source: ArtifactError::MissingArtifact { id } - } if id == artifact_id("tool-result-3") + } if id == artifact_id("tool-result-4") )); } @@ -262,13 +262,13 @@ async fn execute_tool_blank_json_outcome_keeps_pending_without_artifact_or_resul merry_runtime::RuntimeError::UnsupportedToolResultContent { artifact_id, content_kind: ArtifactContentKind::Json - } if artifact_id.as_str() == "tool-result-3" + } if artifact_id.as_str() == "tool-result-4" )); assert_eq!(before, after); assert_eq!(runtime.pending_tool_calls().await, vec![pending]); let evidence_err = runtime .evidence_ref( - &artifact_id("tool-result-3"), + &artifact_id("tool-result-4"), EvidenceLocator::whole_artifact(), ) .await @@ -277,7 +277,7 @@ async fn execute_tool_blank_json_outcome_keeps_pending_without_artifact_or_resul evidence_err, merry_runtime::RuntimeError::Artifact { source: ArtifactError::MissingArtifact { id } - } if id == artifact_id("tool-result-3") + } if id == artifact_id("tool-result-4") )); } @@ -318,7 +318,7 @@ async fn executor_reentrant_runtime_mutations_are_rejected_while_outer_execution ["ArtifactRecorded", "ToolCallResolved"] ); let result = resolved_tool_result(&execution_events); - assert_eq!(result.artifact().id().as_str(), "tool-result-3"); + assert_eq!(result.artifact().id().as_str(), "tool-result-4"); assert_eq!(result.call_id(), pending.id()); assert!(runtime.pending_tool_calls().await.is_empty()); let outer_evidence = runtime diff --git a/crates/merry-runtime/tests/provider_boundary/tool_results.rs b/crates/merry-runtime/tests/provider_boundary/tool_results.rs index f04cf0b4..2f5effb1 100644 --- a/crates/merry-runtime/tests/provider_boundary/tool_results.rs +++ b/crates/merry-runtime/tests/provider_boundary/tool_results.rs @@ -44,7 +44,7 @@ async fn submit_tool_result_records_success_artifact_resolves_pending_and_update .iter() .map(|event| event.sequence) .collect::>(), - vec![3, 4] + vec![4, 5] ); assert!(matches!( &events[0].payload, @@ -76,17 +76,17 @@ async fn submit_tool_result_records_success_artifact_resolves_pending_and_update kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, order: 3, kind: LedgerFactKind::ArtifactRecorded, }, LedgerProjection::Lifecycle { - sequence: 4, + sequence: 5, order: 4, kind: LedgerFactKind::ToolCallResolved, }, @@ -104,7 +104,7 @@ async fn submit_tool_result_records_success_artifact_resolves_pending_and_update .iter() .map(|event| event.sequence) .collect::>(), - vec![5, 6, 7] + vec![6, 7, 8, 9] ); } @@ -240,7 +240,7 @@ async fn artifact_error_while_submitting_tool_result_keeps_call_pending_and_sequ .iter() .map(|event| event.sequence) .collect::>(), - vec![4, 5] + vec![5, 6, 7] ); } @@ -346,7 +346,7 @@ async fn blank_text_tool_result_keeps_call_pending_and_sequence_stable() { .iter() .map(|event| event.sequence) .collect::>(), - vec![3, 4] + vec![4, 5, 6] ); } diff --git a/crates/merry-runtime/tests/provider_boundary/tool_streaming.rs b/crates/merry-runtime/tests/provider_boundary/tool_streaming.rs index 6afc6855..2bea6ad2 100644 --- a/crates/merry-runtime/tests/provider_boundary/tool_streaming.rs +++ b/crates/merry-runtime/tests/provider_boundary/tool_streaming.rs @@ -205,7 +205,7 @@ async fn provider_tool_call_pending_preserves_id_name_arguments_and_ledger_fact( kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, @@ -294,17 +294,17 @@ async fn repeated_provider_tool_call_id_after_pending_fails_without_second_pendi kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 2, + sequence: 3, order: 2, kind: LedgerFactKind::ToolCallPending, }, LedgerProjection::Lifecycle { - sequence: 3, + sequence: 4, order: 3, kind: LedgerFactKind::StepStarted, }, LedgerProjection::Lifecycle { - sequence: 4, + sequence: 6, order: 4, kind: LedgerFactKind::Failed, }, diff --git a/crates/merry-runtime/tests/public_events.rs b/crates/merry-runtime/tests/public_events.rs index f6b803e1..1af73901 100644 --- a/crates/merry-runtime/tests/public_events.rs +++ b/crates/merry-runtime/tests/public_events.rs @@ -91,6 +91,13 @@ async fn collect_journal_stream(runtime: &Runtime, text: &str) -> Vec Vec<&RuntimeEvent> { + events + .iter() + .filter(|event| !matches!(event, RuntimeEvent::ModelOutputRateUpdated { .. })) + .collect() +} + #[tokio::test(flavor = "current_thread")] async fn assistant_output_projects_to_public_assistant_message() { let provider = FakeModelProvider::new(vec![Ok(completed_text_event("hello public event"))]); @@ -100,21 +107,26 @@ async fn assistant_output_projects_to_public_assistant_message() { .expect("runtime should build"); let events = collect_public_stream(&runtime, "Say hello.").await; + let durable = without_rate_updates(&events); - assert!(matches!(events[0], RuntimeEvent::SessionStarted { .. })); - assert!(matches!(events[1], RuntimeEvent::StepStarted { .. })); + assert!(matches!(durable[0], RuntimeEvent::SessionStarted { .. })); + assert!(matches!(durable[1], RuntimeEvent::StepStarted { .. })); let RuntimeEvent::AssistantMessage { text, artifact, source, - } = &events[2] + } = durable[2] else { - panic!("expected assistant message, got {:?}", events[2]); + panic!("expected assistant message, got {:?}", durable[2]); }; assert_eq!(text, "hello public event"); assert_eq!(artifact.kind(), &ArtifactKind::Text); - assert_eq!(source.sequence, 2); - assert!(matches!(events[3], RuntimeEvent::StepCompleted { .. })); + assert_eq!(source.sequence, 3); + assert!(events.iter().any(|event| matches!( + event, + RuntimeEvent::ModelOutputRateUpdated { rate: None, .. } + ))); + assert!(matches!(durable[3], RuntimeEvent::StepCompleted { .. })); } #[tokio::test(flavor = "current_thread")] @@ -134,25 +146,26 @@ async fn streamed_text_delta_projects_to_public_assistant_message_delta() { .expect("runtime should build"); let events = collect_public_stream(&runtime, "Say hello.").await; + let durable = without_rate_updates(&events); - assert!(matches!(events[0], RuntimeEvent::SessionStarted { .. })); - assert!(matches!(events[1], RuntimeEvent::StepStarted { .. })); + assert!(matches!(durable[0], RuntimeEvent::SessionStarted { .. })); + assert!(matches!(durable[1], RuntimeEvent::StepStarted { .. })); assert!(matches!( - &events[2], + durable[2], RuntimeEvent::AssistantMessageDelta { delta, source } - if delta == "hel" && source.sequence == 2 + if delta == "hel" && source.sequence == 3 )); assert!(matches!( - &events[3], + durable[3], RuntimeEvent::AssistantMessageDelta { delta, source } - if delta == "lo" && source.sequence == 3 + if delta == "lo" && source.sequence == 4 )); assert!(matches!( - &events[4], + durable[4], RuntimeEvent::AssistantMessage { text, source, .. } - if text == "hello" && source.sequence == 4 + if text == "hello" && source.sequence == 5 )); - assert!(matches!(events[5], RuntimeEvent::StepCompleted { .. })); + assert!(matches!(durable[5], RuntimeEvent::StepCompleted { .. })); } #[tokio::test(flavor = "current_thread")] @@ -168,13 +181,14 @@ async fn usage_update_projects_to_public_stream_and_getter() { assert_eq!(runtime.usage().await, None); let events = collect_public_stream(&runtime, "Say hello.").await; + let durable = without_rate_updates(&events); - assert!(matches!(events[0], RuntimeEvent::SessionStarted { .. })); - assert!(matches!(events[1], RuntimeEvent::StepStarted { .. })); - let RuntimeEvent::UsageUpdated { usage, source } = &events[2] else { - panic!("expected usage update, got {:?}", events[2]); + assert!(matches!(durable[0], RuntimeEvent::SessionStarted { .. })); + assert!(matches!(durable[1], RuntimeEvent::StepStarted { .. })); + let RuntimeEvent::UsageUpdated { usage, source } = durable[2] else { + panic!("expected usage update, got {:?}", durable[2]); }; - assert_eq!(source.sequence, 2); + assert_eq!(source.sequence, 3); assert_eq!( usage.last, ModelUsage::with_details(12, Some(8), 5, None, 17) @@ -182,8 +196,8 @@ async fn usage_update_projects_to_public_stream_and_getter() { assert_eq!(usage.total, usage.last); assert!(usage.context.is_some()); assert!(usage.compaction.is_some()); - assert!(matches!(events[3], RuntimeEvent::AssistantMessage { .. })); - assert!(matches!(events[4], RuntimeEvent::StepCompleted { .. })); + assert!(matches!(durable[3], RuntimeEvent::AssistantMessage { .. })); + assert!(matches!(durable[4], RuntimeEvent::StepCompleted { .. })); assert_eq!(runtime.usage().await, Some(usage.clone())); } @@ -326,20 +340,34 @@ async fn one_slot_commentary_tool_batch_projects_in_order_after_atomic_state_com events.next().await.expect("step started"), RuntimeEvent::StepStarted { .. } )); + let commentary = loop { + let event = events.next().await.expect("assistant commentary"); + if matches!(event, RuntimeEvent::ModelOutputRateUpdated { .. }) { + continue; + } + break event; + }; assert!(matches!( - events.next().await.expect("assistant commentary"), + commentary, RuntimeEvent::AssistantMessage { text, source, .. } - if text == "Public commentary before tool." && source.sequence == 2 + if text == "Public commentary before tool." && source.sequence == 3 )); assert_eq!( runtime.pending_tool_calls().await.len(), 1, "tool state must already be committed when commentary is visible" ); + let tool_started = loop { + let event = events.next().await.expect("tool started"); + if matches!(event, RuntimeEvent::ModelOutputRateUpdated { .. }) { + continue; + } + break event; + }; assert!(matches!( - events.next().await.expect("tool started"), + tool_started, RuntimeEvent::ToolCallStarted { call, source } - if call.id().as_str() == "call-public-commentary-tool" && source.sequence == 3 + if call.id().as_str() == "call-public-commentary-tool" && source.sequence == 4 )); } @@ -386,8 +414,8 @@ async fn stream_and_journal_stream_are_separate_surfaces() { let journal_events = collect_journal_stream(&runtime, "Raw journal.").await; - assert!(matches!( - journal_events[2].payload, + assert!(journal_events.iter().any(|event| matches!( + event.payload, RuntimeJournalPayload::AssistantOutputRecorded { .. } - )); + ))); } diff --git a/crates/merry-tools/tests/runtime_integration/support.rs b/crates/merry-tools/tests/runtime_integration/support.rs index 0a8322d5..09f5406c 100644 --- a/crates/merry-tools/tests/runtime_integration/support.rs +++ b/crates/merry-tools/tests/runtime_integration/support.rs @@ -319,6 +319,12 @@ fn assert_pending_tool_call_events(events: &[RuntimeJournalEvent]) { pub(super) fn event_kind_names(events: &[RuntimeJournalEvent]) -> Vec<&'static str> { events .iter() + .filter(|event| { + !matches!( + event.payload, + RuntimeJournalPayload::ModelOutputRateUpdated { .. } + ) + }) .map(|event| match event.payload { RuntimeJournalPayload::SessionStarted => "SessionStarted", RuntimeJournalPayload::StepStarted => "StepStarted", diff --git a/crates/merry/src/events.rs b/crates/merry/src/events.rs index fd1913e4..87aea42e 100644 --- a/crates/merry/src/events.rs +++ b/crates/merry/src/events.rs @@ -1,7 +1,7 @@ //! Runtime event protocol types. pub use merry_core::{ - ArtifactId, ArtifactKind, ArtifactRef, PendingToolCall, RuntimeEvent, RuntimeEventSource, + ArtifactId, ArtifactKind, ArtifactRef, ModelOutputRate, OutputTimingQuality, OutputTokenSource, PendingToolCall, RuntimeEvent, RuntimeEventSource, RuntimeJournalEvent, RuntimeJournalPayload, SubagentStatus, ToolCallId, ToolCallResult, ToolCallResultStatus, ToolOutput, }; diff --git a/crates/merry/src/run_result.rs b/crates/merry/src/run_result.rs index c85117bf..13beff53 100644 --- a/crates/merry/src/run_result.rs +++ b/crates/merry/src/run_result.rs @@ -40,7 +40,7 @@ impl RunResult { &self.status } - /// Returns SDK-facing events in durable emission order. + /// Returns retained SDK-facing evidence in emission order, excluding live deltas and rates. #[must_use] pub fn events(&self) -> &[RuntimeEvent] { &self.events @@ -70,7 +70,7 @@ impl RunResult { self.session_usage.as_ref() } - /// Consumes the result and returns its public events. + /// Consumes the result and returns retained public events, excluding live deltas and rates. #[must_use] pub fn into_events(self) -> Vec { self.events diff --git a/crates/merry/tests/agent/events.rs b/crates/merry/tests/agent/events.rs index d049ff77..0cb2c273 100644 --- a/crates/merry/tests/agent/events.rs +++ b/crates/merry/tests/agent/events.rs @@ -31,5 +31,7 @@ async fn stream_result_projects_the_same_public_contract() { assert_eq!(result.status(), &AgentLoopStatus::Completed); assert_eq!(result.final_output(), Some("streamed")); + assert!(events.iter().any(RuntimeEvent::is_transient)); + events.retain(|event| !event.is_transient()); assert_eq!(events, result.events()); } diff --git a/crates/merry/tests/agent/native_tools.rs b/crates/merry/tests/agent/native_tools.rs index fec1cdc3..2cd6986f 100644 --- a/crates/merry/tests/agent/native_tools.rs +++ b/crates/merry/tests/agent/native_tools.rs @@ -145,6 +145,8 @@ async fn typed_profile_tool_is_executed_inside_event_only_stream() { assert_eq!(executions.load(Ordering::SeqCst), 1); assert_eq!(result.status(), &AgentLoopStatus::Completed); assert_eq!(result.final_output(), Some("order streamed")); + assert!(events.iter().any(RuntimeEvent::is_transient)); + events.retain(|event| !event.is_transient()); assert_eq!(events, result.events()); } diff --git a/sdks/python/README.md b/sdks/python/README.md index e5393ab3..9a4aa119 100644 --- a/sdks/python/README.md +++ b/sdks/python/README.md @@ -167,6 +167,12 @@ wait for Rust to stop provider and tool work, so the returned result is durable. The Python async task may be cancelled; the SDK requests Rust cancellation and re-raises `asyncio.CancelledError`. +Text deltas and output-rate observations are live-only events: consume them from +the run stream. `result.events` retains terminal evidence and lifecycle events, +not these transient updates, so long streamed output does not accumulate a +second event-by-event copy in the result. Final assistant text remains available +through `result.final_output` and assistant-message events. + `AgentBuilder` is single-use for `build()` and `resume()`. A native operation that has consumed the builder also makes the Python builder terminal, including when that operation later fails. Start a new builder to retry that operation. diff --git a/sdks/python/merry/__init__.py b/sdks/python/merry/__init__.py index b216f49a..516c8953 100644 --- a/sdks/python/merry/__init__.py +++ b/sdks/python/merry/__init__.py @@ -47,9 +47,13 @@ FinalOutputRecordedPayload, InteractiveRunState, InteractiveRunStateChangedPayload, + ModelOutputRate, + ModelOutputRateUpdatedPayload, ModelRetryAttemptStartedPayload, ModelRetryExhaustedPayload, ModelRetryScheduledPayload, + OutputTimingQuality, + OutputTokenSource, PlanAttemptFinishedPayload, PlanAttemptProgressReportedPayload, PlanDirectiveUpdatedPayload, @@ -149,6 +153,8 @@ "MerryProviderError", "MerryRuntimeError", "MerryToolError", + "ModelOutputRate", + "ModelOutputRateUpdatedPayload", "ModelRetryAttemptStartedPayload", "ModelRetryExhaustedPayload", "ModelRetryScheduledPayload", @@ -156,6 +162,8 @@ "NativeMerryError", "OpenAICompatible", "OpenAICompatibleProvider", + "OutputTimingQuality", + "OutputTokenSource", "PatchConfig", "PlanAttemptFinishedPayload", "PlanAttemptProgressReportedPayload", diff --git a/sdks/python/merry/_event_fields.py b/sdks/python/merry/_event_fields.py new file mode 100644 index 00000000..e5a355b3 --- /dev/null +++ b/sdks/python/merry/_event_fields.py @@ -0,0 +1,76 @@ +"""Strict field contracts for normalized runtime events.""" + +from ._event_types import EventType + + +def _event_fields( + *fields: str, optional: tuple[str, ...] = () +) -> tuple[frozenset[str], frozenset[str]]: + return frozenset(("type", *fields)), frozenset(optional) + + +EVENT_FIELDS: dict[EventType, tuple[frozenset[str], frozenset[str]]] = { + EventType.MODEL_OUTPUT_RATE_UPDATED: _event_fields("rate", "source"), + EventType.SESSION_STARTED: _event_fields("source"), + EventType.STEP_STARTED: _event_fields("source"), + EventType.STEP_COMPLETED: _event_fields("source"), + EventType.COMPACTION_STARTED: _event_fields("source"), + EventType.COMPACTION_COMPLETED: _event_fields( + "checkpoint_id", "covered_history_item_count", "source" + ), + EventType.USAGE_UPDATED: _event_fields("usage", "source"), + EventType.ASSISTANT_MESSAGE: _event_fields("text", "artifact", "source"), + EventType.ASSISTANT_MESSAGE_DELTA: _event_fields("delta", "source"), + EventType.TOOL_CALL_STARTED: _event_fields("call", "source"), + EventType.TOOL_CALL_BATCH_STARTED: _event_fields("batch", "source"), + EventType.TOOL_CALL_FINISHED: _event_fields("result", "output", "source"), + EventType.FINAL_OUTPUT_RECORDED: _event_fields("call_id", "artifact", "source"), + EventType.MODEL_RETRY_ATTEMPT_STARTED: _event_fields( + "attempt", "max_attempts", "source" + ), + EventType.MODEL_RETRY_SCHEDULED: _event_fields( + "attempt", "next_attempt", "max_attempts", "delay_ms", "error_kind", "source" + ), + EventType.MODEL_RETRY_EXHAUSTED: _event_fields( + "attempts_run", "max_attempts", "error_kind", "source" + ), + EventType.EVIDENCE_REFERENCED: _event_fields("evidence", "source"), + EventType.SKILL_USED: _event_fields( + "skill_name", "skill_md_path", "tool_call_id", "artifact", "source" + ), + EventType.SUBAGENT_SPAWNED: _event_fields( + "agent_id", "task_id", "task_anchor", "source" + ), + EventType.SUBAGENT_STARTED: _event_fields("agent_id", "task_id", "source"), + EventType.SUBAGENT_STATUS_CHANGED: _event_fields( + "agent_id", "task_id", "status", "source" + ), + EventType.SUBAGENT_COMPLETED: _event_fields( + "agent_id", "task_id", "summary", "output_paths", "changed_paths", "source" + ), + EventType.SUBAGENT_FAILED: _event_fields( + "agent_id", "task_id", "diagnostic", "source" + ), + EventType.SUBAGENT_CANCELLED: _event_fields( + "agent_id", "task_id", "diagnostic", "source" + ), + EventType.PLAN_UPDATED: _event_fields("snapshot", "summary", "source"), + EventType.PLAN_PHASE_CHANGED: _event_fields("plan_id", "phase", "source"), + EventType.PLAN_NODE_READY: _event_fields( + "plan_id", "node_id", "node_revision", "source" + ), + EventType.PLAN_LEASE_STARTED: _event_fields("lease", "source"), + EventType.PLAN_PROGRESS_UPDATED: _event_fields("progress", "source"), + EventType.PLAN_PROGRESS_REVIEW_REQUESTED: _event_fields( + "plan_id", "attempt_id", "reason", "source" + ), + EventType.PLAN_ATTEMPT_PROGRESS_REPORTED: _event_fields("progress", "source"), + EventType.PLAN_DIRECTIVE_UPDATED: _event_fields("directive", "source"), + EventType.PLAN_ATTEMPT_FINISHED: _event_fields("attempt", "source"), + EventType.RUN_FAILED: _event_fields("diagnostic", "source"), + EventType.RUN_CANCELLED: _event_fields("diagnostic", "source"), + EventType.INTERACTIVE_RUN_STATE_CHANGED: _event_fields("state"), + EventType.QUEUED_INPUT_ACCEPTED: _event_fields("lane", "inputs"), + EventType.QUEUED_INPUTS_CHANGED: _event_fields("inputs"), + EventType.CLOSED: _event_fields(), +} diff --git a/sdks/python/merry/_event_lifecycle.py b/sdks/python/merry/_event_lifecycle.py new file mode 100644 index 00000000..497e5d5d --- /dev/null +++ b/sdks/python/merry/_event_lifecycle.py @@ -0,0 +1,137 @@ +"""Typed session, model-attempt, usage, and terminal lifecycle payloads.""" + +from dataclasses import dataclass +from typing import ClassVar + +from ._event_types import ( + EventDiagnostic, + EventSource, + EventType, + SourcedEventPayload, + _validate_event_source, + _validate_identifier, + _validate_nonnegative, +) +from ._models import SessionUsage + + +@dataclass(frozen=True, slots=True) +class SessionStartedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.SESSION_STARTED + source: EventSource + + +@dataclass(frozen=True, slots=True) +class StepStartedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.STEP_STARTED + source: EventSource + + +@dataclass(frozen=True, slots=True) +class StepCompletedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.STEP_COMPLETED + source: EventSource + + +@dataclass(frozen=True, slots=True) +class CompactionStartedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.COMPACTION_STARTED + source: EventSource + + +@dataclass(frozen=True, slots=True) +class CompactionCompletedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.COMPACTION_COMPLETED + checkpoint_id: str + covered_history_item_count: int + source: EventSource + + def __post_init__(self) -> None: + _validate_event_source(self.source) + _validate_identifier("checkpoint id", self.checkpoint_id, 256) + _validate_nonnegative( + "covered history item count", self.covered_history_item_count + ) + + +@dataclass(frozen=True, slots=True) +class UsageUpdatedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.USAGE_UPDATED + usage: SessionUsage + source: EventSource + + def __post_init__(self) -> None: + _validate_event_source(self.source) + if not isinstance(self.usage, SessionUsage): + raise TypeError("usage must be a SessionUsage") + + +@dataclass(frozen=True, slots=True) +class ModelRetryAttemptStartedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.MODEL_RETRY_ATTEMPT_STARTED + attempt: int + max_attempts: int + source: EventSource + + def __post_init__(self) -> None: + _validate_event_source(self.source) + _validate_nonnegative("retry attempt", self.attempt) + _validate_nonnegative("maximum retry attempts", self.max_attempts) + + +@dataclass(frozen=True, slots=True) +class ModelRetryScheduledPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.MODEL_RETRY_SCHEDULED + attempt: int + next_attempt: int + max_attempts: int + delay_ms: int + error_kind: str + source: EventSource + + def __post_init__(self) -> None: + _validate_event_source(self.source) + _validate_nonnegative("retry attempt", self.attempt) + _validate_nonnegative("next retry attempt", self.next_attempt) + _validate_nonnegative("maximum retry attempts", self.max_attempts) + _validate_nonnegative("retry delay", self.delay_ms) + _validate_identifier("retry error kind", self.error_kind, 256) + + +@dataclass(frozen=True, slots=True) +class ModelRetryExhaustedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.MODEL_RETRY_EXHAUSTED + attempts_run: int + max_attempts: int + error_kind: str + source: EventSource + + def __post_init__(self) -> None: + _validate_event_source(self.source) + _validate_nonnegative("retry attempts run", self.attempts_run) + _validate_nonnegative("maximum retry attempts", self.max_attempts) + _validate_identifier("retry error kind", self.error_kind, 256) + + +@dataclass(frozen=True, slots=True) +class RunFailedPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.RUN_FAILED + diagnostic: EventDiagnostic + source: EventSource + + def __post_init__(self) -> None: + _validate_event_source(self.source) + if not isinstance(self.diagnostic, EventDiagnostic): + raise TypeError("run diagnostic must be an EventDiagnostic") + + +@dataclass(frozen=True, slots=True) +class RunCancelledPayload(SourcedEventPayload): + event_type: ClassVar[EventType] = EventType.RUN_CANCELLED + diagnostic: EventDiagnostic + source: EventSource + + def __post_init__(self) -> None: + _validate_event_source(self.source) + if not isinstance(self.diagnostic, EventDiagnostic): + raise TypeError("run diagnostic must be an EventDiagnostic") diff --git a/sdks/python/merry/_event_parser.py b/sdks/python/merry/_event_parser.py index 23f91b32..18f8c16d 100644 --- a/sdks/python/merry/_event_parser.py +++ b/sdks/python/merry/_event_parser.py @@ -3,9 +3,8 @@ from __future__ import annotations from collections.abc import Mapping -from enum import Enum -from typing import TypeVar +from ._event_fields import EVENT_FIELDS from ._event_payloads import ( AssistantMessageDeltaPayload, AssistantMessagePayload, @@ -16,6 +15,7 @@ EvidenceReferencedPayload, FinalOutputRecordedPayload, InteractiveRunStateChangedPayload, + ModelOutputRateUpdatedPayload, ModelRetryAttemptStartedPayload, ModelRetryExhaustedPayload, ModelRetryScheduledPayload, @@ -67,6 +67,14 @@ RuntimeToolResultStatus, SubagentStatus, ) +from ._event_values import ( + _enum_value, + _event_text, + _optional_string, + _required_int, + _required_string, + _required_strings, +) from ._events import Event from ._json import ( JsonObject, @@ -75,80 +83,7 @@ validate_object_keys, ) from ._models import SessionUsage - -EnumT = TypeVar("EnumT", bound=Enum) - - -def _event_fields( - *fields: str, optional: tuple[str, ...] = () -) -> tuple[frozenset[str], frozenset[str]]: - return frozenset(("type", *fields)), frozenset(optional) - - -_EVENT_FIELDS: dict[EventType, tuple[frozenset[str], frozenset[str]]] = { - EventType.SESSION_STARTED: _event_fields("source"), - EventType.STEP_STARTED: _event_fields("source"), - EventType.STEP_COMPLETED: _event_fields("source"), - EventType.COMPACTION_STARTED: _event_fields("source"), - EventType.COMPACTION_COMPLETED: _event_fields( - "checkpoint_id", "covered_history_item_count", "source" - ), - EventType.USAGE_UPDATED: _event_fields("usage", "source"), - EventType.ASSISTANT_MESSAGE: _event_fields("text", "artifact", "source"), - EventType.ASSISTANT_MESSAGE_DELTA: _event_fields("delta", "source"), - EventType.TOOL_CALL_STARTED: _event_fields("call", "source"), - EventType.TOOL_CALL_BATCH_STARTED: _event_fields("batch", "source"), - EventType.TOOL_CALL_FINISHED: _event_fields("result", "output", "source"), - EventType.FINAL_OUTPUT_RECORDED: _event_fields("call_id", "artifact", "source"), - EventType.MODEL_RETRY_ATTEMPT_STARTED: _event_fields( - "attempt", "max_attempts", "source" - ), - EventType.MODEL_RETRY_SCHEDULED: _event_fields( - "attempt", "next_attempt", "max_attempts", "delay_ms", "error_kind", "source" - ), - EventType.MODEL_RETRY_EXHAUSTED: _event_fields( - "attempts_run", "max_attempts", "error_kind", "source" - ), - EventType.EVIDENCE_REFERENCED: _event_fields("evidence", "source"), - EventType.SKILL_USED: _event_fields( - "skill_name", "skill_md_path", "tool_call_id", "artifact", "source" - ), - EventType.SUBAGENT_SPAWNED: _event_fields( - "agent_id", "task_id", "task_anchor", "source" - ), - EventType.SUBAGENT_STARTED: _event_fields("agent_id", "task_id", "source"), - EventType.SUBAGENT_STATUS_CHANGED: _event_fields( - "agent_id", "task_id", "status", "source" - ), - EventType.SUBAGENT_COMPLETED: _event_fields( - "agent_id", "task_id", "summary", "output_paths", "changed_paths", "source" - ), - EventType.SUBAGENT_FAILED: _event_fields( - "agent_id", "task_id", "diagnostic", "source" - ), - EventType.SUBAGENT_CANCELLED: _event_fields( - "agent_id", "task_id", "diagnostic", "source" - ), - EventType.PLAN_UPDATED: _event_fields("snapshot", "summary", "source"), - EventType.PLAN_PHASE_CHANGED: _event_fields("plan_id", "phase", "source"), - EventType.PLAN_NODE_READY: _event_fields( - "plan_id", "node_id", "node_revision", "source" - ), - EventType.PLAN_LEASE_STARTED: _event_fields("lease", "source"), - EventType.PLAN_PROGRESS_UPDATED: _event_fields("progress", "source"), - EventType.PLAN_PROGRESS_REVIEW_REQUESTED: _event_fields( - "plan_id", "attempt_id", "reason", "source" - ), - EventType.PLAN_ATTEMPT_PROGRESS_REPORTED: _event_fields("progress", "source"), - EventType.PLAN_DIRECTIVE_UPDATED: _event_fields("directive", "source"), - EventType.PLAN_ATTEMPT_FINISHED: _event_fields("attempt", "source"), - EventType.RUN_FAILED: _event_fields("diagnostic", "source"), - EventType.RUN_CANCELLED: _event_fields("diagnostic", "source"), - EventType.INTERACTIVE_RUN_STATE_CHANGED: _event_fields("state"), - EventType.QUEUED_INPUT_ACCEPTED: _event_fields("lane", "inputs"), - EventType.QUEUED_INPUTS_CHANGED: _event_fields("inputs"), - EventType.CLOSED: _event_fields(), -} +from ._output_rate import parse_output_rate def parse_event(value: object) -> Event: @@ -163,7 +98,7 @@ def parse_event(value: object) -> Event: EventType.UNKNOWN, UnknownEventPayload(raw_type, RawEventData(data)) ) - required, optional = _EVENT_FIELDS[event_type] + required, optional = EVENT_FIELDS[event_type] validate_object_keys( data, "native runtime event", @@ -194,6 +129,10 @@ def _parse_payload( ) case EventType.USAGE_UPDATED: return UsageUpdatedPayload(_parse_usage(data["usage"]), _source(data)) + case EventType.MODEL_OUTPUT_RATE_UPDATED: + return ModelOutputRateUpdatedPayload( + parse_output_rate(data["rate"]), _source(data) + ) case EventType.ASSISTANT_MESSAGE: return AssistantMessagePayload( _event_text(data, "text"), @@ -508,54 +447,3 @@ def _parse_usage(value: object) -> SessionUsage: from ._models import parse_session_usage return parse_session_usage(value) - - -def _enum_value(enum_type: type[EnumT], value: object, label: str) -> EnumT: - if not isinstance(value, str): - raise TypeError(f"{label} must be a string") - try: - return enum_type(value) - except ValueError as error: - raise TypeError(f"{label} is unsupported") from error - - -def _event_text(data: Mapping[str, JsonValue], key: str) -> str: - value = data[key] - if not isinstance(value, str): - raise TypeError(f"event field {key!r} must be a string") - return value - - -def _required_string(data: Mapping[str, JsonValue], key: str) -> str: - value = data[key] - if not isinstance(value, str): - raise TypeError(f"event field {key!r} must be a string") - if not value.strip(): - raise ValueError(f"event field {key!r} must not be blank") - return value - - -def _optional_string(value: object) -> str | None: - if value is None: - return None - if not isinstance(value, str): - raise TypeError("optional event field must be a string or null") - return value - - -def _required_int(data: Mapping[str, JsonValue], key: str) -> int: - value = data[key] - if isinstance(value, bool) or not isinstance(value, int): - raise TypeError(f"event field {key!r} must be an integer") - return value - - -def _required_strings(value: object, label: str) -> tuple[str, ...]: - if not isinstance(value, list): - raise TypeError(f"{label} must be a list") - values: list[str] = [] - for item in value: - if not isinstance(item, str): - raise TypeError(f"{label} must contain only strings") - values.append(item) - return tuple(values) diff --git a/sdks/python/merry/_event_payloads.py b/sdks/python/merry/_event_payloads.py index a789c837..930634dc 100644 --- a/sdks/python/merry/_event_payloads.py +++ b/sdks/python/merry/_event_payloads.py @@ -5,6 +5,19 @@ from dataclasses import dataclass from typing import ClassVar, TypeAlias +from ._event_lifecycle import ( + CompactionCompletedPayload, + CompactionStartedPayload, + ModelRetryAttemptStartedPayload, + ModelRetryExhaustedPayload, + ModelRetryScheduledPayload, + RunCancelledPayload, + RunFailedPayload, + SessionStartedPayload, + StepCompletedPayload, + StepStartedPayload, + UsageUpdatedPayload, +) from ._event_types import ( ArtifactReference, EventDiagnostic, @@ -28,58 +41,7 @@ _validate_nonnegative, _validate_text, ) -from ._models import SessionUsage - - -@dataclass(frozen=True, slots=True) -class SessionStartedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.SESSION_STARTED - source: EventSource - - -@dataclass(frozen=True, slots=True) -class StepStartedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.STEP_STARTED - source: EventSource - - -@dataclass(frozen=True, slots=True) -class StepCompletedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.STEP_COMPLETED - source: EventSource - - -@dataclass(frozen=True, slots=True) -class CompactionStartedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.COMPACTION_STARTED - source: EventSource - - -@dataclass(frozen=True, slots=True) -class CompactionCompletedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.COMPACTION_COMPLETED - checkpoint_id: str - covered_history_item_count: int - source: EventSource - - def __post_init__(self) -> None: - _validate_event_source(self.source) - _validate_identifier("checkpoint id", self.checkpoint_id, 256) - _validate_nonnegative( - "covered history item count", self.covered_history_item_count - ) - - -@dataclass(frozen=True, slots=True) -class UsageUpdatedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.USAGE_UPDATED - usage: SessionUsage - source: EventSource - - def __post_init__(self) -> None: - _validate_event_source(self.source) - if not isinstance(self.usage, SessionUsage): - raise TypeError("usage must be a SessionUsage") +from ._output_rate import ModelOutputRateUpdatedPayload @dataclass(frozen=True, slots=True) @@ -172,53 +134,6 @@ def __post_init__(self) -> None: raise TypeError("final output artifact must be an ArtifactReference") -@dataclass(frozen=True, slots=True) -class ModelRetryAttemptStartedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.MODEL_RETRY_ATTEMPT_STARTED - attempt: int - max_attempts: int - source: EventSource - - def __post_init__(self) -> None: - _validate_event_source(self.source) - _validate_nonnegative("retry attempt", self.attempt) - _validate_nonnegative("maximum retry attempts", self.max_attempts) - - -@dataclass(frozen=True, slots=True) -class ModelRetryScheduledPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.MODEL_RETRY_SCHEDULED - attempt: int - next_attempt: int - max_attempts: int - delay_ms: int - error_kind: str - source: EventSource - - def __post_init__(self) -> None: - _validate_event_source(self.source) - _validate_nonnegative("retry attempt", self.attempt) - _validate_nonnegative("next retry attempt", self.next_attempt) - _validate_nonnegative("maximum retry attempts", self.max_attempts) - _validate_nonnegative("retry delay", self.delay_ms) - _validate_identifier("retry error kind", self.error_kind, 256) - - -@dataclass(frozen=True, slots=True) -class ModelRetryExhaustedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.MODEL_RETRY_EXHAUSTED - attempts_run: int - max_attempts: int - error_kind: str - source: EventSource - - def __post_init__(self) -> None: - _validate_event_source(self.source) - _validate_nonnegative("retry attempts run", self.attempts_run) - _validate_nonnegative("maximum retry attempts", self.max_attempts) - _validate_identifier("retry error kind", self.error_kind, 256) - - @dataclass(frozen=True, slots=True) class EvidenceReferencedPayload(SourcedEventPayload): event_type: ClassVar[EventType] = EventType.EVIDENCE_REFERENCED @@ -462,30 +377,6 @@ def __post_init__(self) -> None: raise TypeError("plan attempt must be RawEventData") -@dataclass(frozen=True, slots=True) -class RunFailedPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.RUN_FAILED - diagnostic: EventDiagnostic - source: EventSource - - def __post_init__(self) -> None: - _validate_event_source(self.source) - if not isinstance(self.diagnostic, EventDiagnostic): - raise TypeError("run diagnostic must be an EventDiagnostic") - - -@dataclass(frozen=True, slots=True) -class RunCancelledPayload(SourcedEventPayload): - event_type: ClassVar[EventType] = EventType.RUN_CANCELLED - diagnostic: EventDiagnostic - source: EventSource - - def __post_init__(self) -> None: - _validate_event_source(self.source) - if not isinstance(self.diagnostic, EventDiagnostic): - raise TypeError("run diagnostic must be an EventDiagnostic") - - @dataclass(frozen=True, slots=True) class InteractiveRunStateChangedPayload(EventPayload): event_type: ClassVar[EventType] = EventType.INTERACTIVE_RUN_STATE_CHANGED @@ -556,6 +447,7 @@ def _validate_strings(name: str, values: tuple[str, ...], maximum: int) -> None: | CompactionStartedPayload | CompactionCompletedPayload | UsageUpdatedPayload + | ModelOutputRateUpdatedPayload | AssistantMessagePayload | AssistantMessageDeltaPayload | ToolCallStartedPayload diff --git a/sdks/python/merry/_event_types.py b/sdks/python/merry/_event_types.py index 46576675..4211756c 100644 --- a/sdks/python/merry/_event_types.py +++ b/sdks/python/merry/_event_types.py @@ -24,6 +24,7 @@ class EventType(str, Enum): COMPACTION_STARTED = "compaction_started" COMPACTION_COMPLETED = "compaction_completed" USAGE_UPDATED = "usage_updated" + MODEL_OUTPUT_RATE_UPDATED = "model_output_rate_updated" ASSISTANT_MESSAGE = "assistant_message" ASSISTANT_MESSAGE_DELTA = "assistant_message_delta" TOOL_CALL_STARTED = "tool_call_started" diff --git a/sdks/python/merry/_event_values.py b/sdks/python/merry/_event_values.py new file mode 100644 index 00000000..6afceca3 --- /dev/null +++ b/sdks/python/merry/_event_values.py @@ -0,0 +1,60 @@ +"""Small validated primitives shared by event JSON decoders.""" + +from collections.abc import Mapping +from enum import Enum +from typing import TypeVar + +from ._json import JsonValue + +EnumT = TypeVar("EnumT", bound=Enum) + + +def _enum_value(enum_type: type[EnumT], value: object, label: str) -> EnumT: + if not isinstance(value, str): + raise TypeError(f"{label} must be a string") + try: + return enum_type(value) + except ValueError as error: + raise TypeError(f"{label} is unsupported") from error + + +def _event_text(data: Mapping[str, JsonValue], key: str) -> str: + value = data[key] + if not isinstance(value, str): + raise TypeError(f"event field {key!r} must be a string") + return value + + +def _required_string(data: Mapping[str, JsonValue], key: str) -> str: + value = data[key] + if not isinstance(value, str): + raise TypeError(f"event field {key!r} must be a string") + if not value.strip(): + raise ValueError(f"event field {key!r} must not be blank") + return value + + +def _optional_string(value: object) -> str | None: + if value is None: + return None + if not isinstance(value, str): + raise TypeError("optional event field must be a string or null") + return value + + +def _required_int(data: Mapping[str, JsonValue], key: str) -> int: + value = data[key] + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError(f"event field {key!r} must be an integer") + return value + + +def _required_strings(value: object, label: str) -> tuple[str, ...]: + if not isinstance(value, list): + raise TypeError(f"{label} must be a list") + values: list[str] = [] + for item in value: + if not isinstance(item, str): + raise TypeError(f"{label} must contain only strings") + values.append(item) + return tuple(values) diff --git a/sdks/python/merry/_events.py b/sdks/python/merry/_events.py index 7222b2f6..9c17caf7 100644 --- a/sdks/python/merry/_events.py +++ b/sdks/python/merry/_events.py @@ -14,6 +14,7 @@ EvidenceReferencedPayload, FinalOutputRecordedPayload, InteractiveRunStateChangedPayload, + ModelOutputRateUpdatedPayload, ModelRetryAttemptStartedPayload, ModelRetryExhaustedPayload, ModelRetryScheduledPayload, @@ -66,6 +67,7 @@ RuntimeToolResultStatus, SubagentStatus, ) +from ._output_rate import ModelOutputRate, OutputTimingQuality, OutputTokenSource @dataclass(frozen=True, slots=True) @@ -102,9 +104,13 @@ def __post_init__(self) -> None: "FinalOutputRecordedPayload", "InteractiveRunState", "InteractiveRunStateChangedPayload", + "ModelOutputRate", + "ModelOutputRateUpdatedPayload", "ModelRetryAttemptStartedPayload", "ModelRetryExhaustedPayload", "ModelRetryScheduledPayload", + "OutputTimingQuality", + "OutputTokenSource", "PlanAttemptFinishedPayload", "PlanAttemptProgressReportedPayload", "PlanDirectiveUpdatedPayload", diff --git a/sdks/python/merry/_output_rate.py b/sdks/python/merry/_output_rate.py new file mode 100644 index 00000000..c76f8695 --- /dev/null +++ b/sdks/python/merry/_output_rate.py @@ -0,0 +1,89 @@ +"""Typed receive-side throughput observations owned by the Rust runtime.""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import Enum +from typing import ClassVar + +from ._event_types import ( + EventSource, + EventType, + SourcedEventPayload, + _validate_event_source, + _validate_nonnegative, +) +from ._event_values import _enum_value, _required_int +from ._json import JsonValue, require_json_object, validate_object_keys + + +class OutputTokenSource(str, Enum): + """Source of the output-token numerator.""" + + ESTIMATED = "estimated" + PROVIDER_USAGE = "provider_usage" + + +class OutputTimingQuality(str, Enum): + """Limitations of the client receive interval, not a server decode clock.""" + + RECEIVE_WINDOW = "receive_window" + PARTIAL_OUTPUT = "partial_output" + CONSUMER_LIMITED = "consumer_limited" + + +@dataclass(frozen=True, slots=True) +class ModelOutputRate: + """Output tokens over the provider's first-to-last receive interval. + + Zero duration cannot yield throughput. Limited timing retains an approximate + receive rate; consumer-limited intervals can include consumer stalls. + Provider usage does not imply that hidden reasoning timing was observed. + """ + + output_tokens: int + elapsed_nanos: int + token_source: OutputTokenSource + timing_quality: OutputTimingQuality + + def __post_init__(self) -> None: + _validate_nonnegative("output tokens", self.output_tokens) + _validate_nonnegative("output elapsed nanoseconds", self.elapsed_nanos) + if not isinstance(self.token_source, OutputTokenSource): + raise TypeError("token_source must be an OutputTokenSource") + if not isinstance(self.timing_quality, OutputTimingQuality): + raise TypeError("timing_quality must be an OutputTimingQuality") + + +@dataclass(frozen=True, slots=True) +class ModelOutputRateUpdatedPayload(SourcedEventPayload): + """Receive-side sample, or None when the runtime clears the active sample.""" + + event_type: ClassVar[EventType] = EventType.MODEL_OUTPUT_RATE_UPDATED + rate: ModelOutputRate | None + source: EventSource + + def __post_init__(self) -> None: + _validate_event_source(self.source) + if self.rate is not None and not isinstance(self.rate, ModelOutputRate): + raise TypeError("rate must be a ModelOutputRate or None") + + +def parse_output_rate(value: JsonValue) -> ModelOutputRate | None: + """Decode a strict runtime observation without recalculating Rust-owned metrics.""" + if value is None: + return None + rate = require_json_object(value, "model output rate") + validate_object_keys( + rate, + "model output rate", + required={"output_tokens", "elapsed_nanos", "token_source", "timing_quality"}, + ) + return ModelOutputRate( + _required_int(rate, "output_tokens"), + _required_int(rate, "elapsed_nanos"), + _enum_value(OutputTokenSource, rate["token_source"], "output token source"), + _enum_value( + OutputTimingQuality, rate["timing_quality"], "output timing quality" + ), + ) diff --git a/sdks/python/merry/_run.py b/sdks/python/merry/_run.py index 302fbef1..9f9131dc 100644 --- a/sdks/python/merry/_run.py +++ b/sdks/python/merry/_run.py @@ -101,7 +101,7 @@ async def next(self) -> Event | ToolCallBatch[OutputT] | None: raise async def result(self) -> RunResult[OutputT]: - """Return the durable terminal result after the run reaches EOF.""" + """Return terminal evidence after EOF, excluding live text deltas and rates.""" if self._result is not None: return self._result diff --git a/sdks/python/tests/test_output_rate.py b/sdks/python/tests/test_output_rate.py new file mode 100644 index 00000000..dfa19e9d --- /dev/null +++ b/sdks/python/tests/test_output_rate.py @@ -0,0 +1,67 @@ +from __future__ import annotations + +import pytest + +import merry +from merry._event_parser import parse_event +from merry._json import JsonObject, JsonValue + + +def rate_event() -> JsonObject: + return { + "type": "model_output_rate_updated", + "rate": { + "output_tokens": 2400, + "elapsed_nanos": 8_000_000_000, + "token_source": "provider_usage", + "timing_quality": "receive_window", + }, + "source": {"session_id": "session-1", "sequence": 3}, + } + + +def test_output_rate_event_exposes_typed_provider_receive_timing() -> None: + event = parse_event(rate_event()) + assert event.type is merry.EventType.MODEL_OUTPUT_RATE_UPDATED + assert isinstance(event.payload, merry.ModelOutputRateUpdatedPayload) + assert event.payload.rate == merry.ModelOutputRate( + 2400, + 8_000_000_000, + merry.OutputTokenSource.PROVIDER_USAGE, + merry.OutputTimingQuality.RECEIVE_WINDOW, + ) + assert event.payload.source.sequence == 3 + + +def test_output_rate_reset_does_not_conflate_missing_with_zero_usage() -> None: + data = rate_event() + data["rate"] = None + event = parse_event(data) + assert isinstance(event.payload, merry.ModelOutputRateUpdatedPayload) + assert event.payload.rate is None + + +@pytest.mark.parametrize( + ("field", "value"), + [ + ("output_tokens", -1), + ("elapsed_nanos", True), + ("token_source", False), + ("timing_quality", "unsupported"), + ("extra", 1), + ], +) +def test_output_rate_rejects_invalid_or_unknown_fields( + field: str, value: JsonValue +) -> None: + data = rate_event() + rate: JsonObject = { + "output_tokens": 2400, + "elapsed_nanos": 8_000_000_000, + "token_source": "provider_usage", + "timing_quality": "receive_window", + } + rate[field] = value + data["rate"] = rate + with pytest.raises((TypeError, ValueError)): + parse_event(data) diff --git a/sdks/python/tests/test_run.py b/sdks/python/tests/test_run.py index 56c942d4..85670557 100644 --- a/sdks/python/tests/test_run.py +++ b/sdks/python/tests/test_run.py @@ -111,26 +111,29 @@ async def scenario() -> merry.RunResult[BaseModel]: def test_pure_newline_delta_keeps_run_active() -> None: - async def scenario() -> merry.RunResult[BaseModel]: + async def scenario() -> tuple[merry.RunResult[BaseModel], list[str]]: run = streamed_text_delta_agent( delta="\r\n", final_text="Order loaded.", session_id="newline-delta", ).stream("Stream a result.") - while await run.next() is not None: - pass - return await run.result() + deltas: list[str] = [] + while (event := await run.next()) is not None: + if isinstance(event, merry.Event) and isinstance( + event.payload, merry.AssistantMessageDeltaPayload + ): + deltas.append(event.payload.delta) + return await run.result(), deltas - result = asyncio.run(scenario()) - - delta_payloads: list[merry.AssistantMessageDeltaPayload] = [] - for event in result.events: - if isinstance(event.payload, merry.AssistantMessageDeltaPayload): - delta_payloads.append(event.payload) + result, deltas = asyncio.run(scenario()) assert result.status is merry.RunStatus.COMPLETED assert result.final_output == "Order loaded." - assert [payload.delta for payload in delta_payloads] == ["\r\n"] + assert deltas == ["\r\n"] + assert not any( + isinstance(event.payload, merry.AssistantMessageDeltaPayload) + for event in result.events + ) def test_result_requires_eof_and_cancel_is_idempotent() -> None: @@ -365,7 +368,11 @@ def test_direct_next_cancellation_persists_a_terminal_result() -> None: async def scenario() -> merry.RunResult[BaseModel]: native = pending_native_agent("direct-next-cancel") run = merry.Agent._from_native(native).stream("Wait for cancellation.") - for expected_type in ("session_started", "step_started"): + for expected_type in ( + "session_started", + "step_started", + "model_output_rate_updated", + ): message = await run.next() assert isinstance(message, merry.Event) assert message.type.value == expected_type @@ -388,7 +395,11 @@ def test_concurrent_next_is_rejected_without_disrupting_active_run() -> None: async def scenario() -> merry.RunResult[BaseModel]: native = pending_native_agent("concurrent-next") run = merry.Agent._from_native(native).stream("Wait for cancellation.") - for expected_type in ("session_started", "step_started"): + for expected_type in ( + "session_started", + "step_started", + "model_output_rate_updated", + ): message = await run.next() assert isinstance(message, merry.Event) assert message.type.value == expected_type