Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 45 additions & 18 deletions crates/core/src/stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,9 @@ use crate::api::runtime::{LlmSanitizeResponseContext, LlmSanitizeResponseFn};
use crate::api::shared::{
metadata_with_otel_error, metadata_with_otel_status, snapshot_event_sanitizers,
};
use crate::codec::response::{AnnotatedLlmResponse, attach_estimated_cost_for_provider};
use crate::codec::response::{
AnnotatedLlmResponse, FinishReason, attach_estimated_cost_for_provider,
};
use crate::codec::traits::LlmResponseCodec;
use crate::error::{FlowError, Result};
use crate::json::Json;
Expand Down Expand Up @@ -86,6 +88,19 @@ pub struct LlmStreamWrapper {
terminal_result: Option<Result<Json>>,
}

#[derive(Clone, Copy, PartialEq, Eq)]
enum StreamTermination {
Complete,
Failed,
Dropped,
}

impl StreamTermination {
const fn is_interrupted(self) -> bool {
!matches!(self, Self::Complete)
}
}

impl LlmStreamWrapper {
/// Create a new `LlmStreamWrapper` around the given raw stream.
///
Expand Down Expand Up @@ -196,33 +211,28 @@ impl LlmStreamWrapper {
self.handle
.optimization_recorder
.close_for_finalization(None);
self.finalization = self.emit_end_event(metadata, true, background_thread);
self.finalization =
self.emit_end_event(metadata, StreamTermination::Dropped, background_thread);
}

fn finish_with_status(
&mut self,
status_code: &'static str,
status_message: Option<String>,
interrupted: bool,
) {
fn finish_cleanly(&mut self) {
if self.ended {
return;
}
self.ended = true;
self.inner.terminalize();
let metadata =
metadata_with_otel_status(self.metadata.clone(), status_code, status_message);
self.finalization = self.emit_end_event(metadata, interrupted, false);
let metadata = metadata_with_otel_status(self.metadata.clone(), "OK", None);
self.finalization = self.emit_end_event(metadata, StreamTermination::Complete, false);
}

fn finish_with_error(&mut self, error: &FlowError, interrupted: bool) {
fn finish_with_error(&mut self, error: &FlowError) {
if self.ended {
return;
}
self.ended = true;
self.inner.terminalize();
let metadata = metadata_with_otel_error(self.metadata.clone(), error);
self.finalization = self.emit_end_event(metadata, interrupted, false);
self.finalization = self.emit_end_event(metadata, StreamTermination::Failed, false);
}

/// Emit the LLM END event with aggregated response data.
Expand All @@ -232,7 +242,7 @@ impl LlmStreamWrapper {
fn emit_end_event(
&mut self,
metadata: Option<Json>,
interrupted: bool,
termination: StreamTermination,
background_thread: bool,
) -> Option<tokio::task::JoinHandle<()>> {
// The finalizer below runs on the caller's Tokio runtime. Register a
Expand Down Expand Up @@ -286,7 +296,14 @@ impl LlmStreamWrapper {
})
})
.flatten();
let interruption = (interrupted
let metadata = if termination == StreamTermination::Dropped
&& has_authoritative_terminal_outcome(annotated_response.as_ref())
{
metadata_with_otel_status(metadata, "OK", None)
} else {
metadata
};
let interruption = (termination.is_interrupted()
&& !has_authoritative_final_usage(annotated_response.as_ref()))
.then_some("stream_interrupted");
Comment thread
coderabbitai[bot] marked this conversation as resolved.
handle
Expand Down Expand Up @@ -460,19 +477,19 @@ impl Stream for LlmStreamWrapper {
match (this.collector)(raw_chunk.clone()) {
Ok(()) => Poll::Ready(Some(Ok(raw_chunk))),
Err(e) => {
this.finish_with_error(&e, true);
this.finish_with_error(&e);
this.terminal_result = Some(Err(e));
self.poll_next(cx)
}
}
}
Poll::Ready(Some(Err(e))) => {
this.finish_with_error(&e, true);
this.finish_with_error(&e);
this.terminal_result = Some(Err(e));
self.poll_next(cx)
}
Poll::Ready(None) => {
this.finish_with_status("OK", None, false);
this.finish_cleanly();
self.poll_next(cx)
}
Poll::Pending => Poll::Pending,
Expand Down Expand Up @@ -513,6 +530,16 @@ fn has_authoritative_final_usage(response: Option<&AnnotatedLlmResponse>) -> boo
})
}

fn has_authoritative_terminal_outcome(response: Option<&AnnotatedLlmResponse>) -> bool {
has_authoritative_final_usage(response)
&& response.is_some_and(|response| {
response
.finish_reason
.as_ref()
.is_some_and(|reason| !matches!(reason, FinishReason::Unknown(_)))
})
}

fn llm_chunk_mark_data(chunk_index: u64, raw_chunk: &Json) -> Json {
if let Some(data) = summarize_openai_chat_chunk(chunk_index, raw_chunk) {
return data;
Expand Down
68 changes: 68 additions & 0 deletions crates/core/tests/integration/stream_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,9 @@ use nemo_relay::api::optimization::LlmOptimizationRecorder;
use nemo_relay::api::runtime::global_context;
use nemo_relay::api::runtime::{LlmJsonStream, LlmStreamInner, NemoRelayContextState};
use nemo_relay::api::subscriber::{deregister_subscriber, flush_subscribers, register_subscriber};
use nemo_relay::codec::openai_responses::{OpenAIResponsesCodec, OpenAIResponsesStreamingCodec};
use nemo_relay::codec::optimization::LlmOptimizationContribution;
use nemo_relay::codec::streaming::StreamingCodec;
use nemo_relay::error::FlowError;
use nemo_relay::error::Result;
use nemo_relay::json::Json;
Expand Down Expand Up @@ -471,6 +473,72 @@ async fn test_stream_wrapper_drop_emits_end_event_for_partial_stream() {
deregister_subscriber("stream_drop_end_test").unwrap();
}

#[tokio::test]
async fn dropped_stream_after_terminal_response_emits_success() {
let _lock = TEST_MUTEX.lock().unwrap();
reset_global();

let events = Arc::new(Mutex::new(Vec::new()));
let captured = events.clone();
register_subscriber(
"stream_terminal_drop_status_test",
Arc::new(move |event: &Event| captured.lock().unwrap().push(event.clone())),
)
.unwrap();

let terminal_event = json!({
"type": "response.completed",
"response": {
"id": "resp_complete",
"model": "gpt-5",
"status": "completed",
"output": [{
"type": "message",
"content": [{"type": "output_text", "text": "done"}]
}],
"usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12}
}
});
let streaming_codec = OpenAIResponsesStreamingCodec::new();
let collector = streaming_codec.collector();
let finalizer = streaming_codec.finalizer();
let request = LlmRequest {
headers: serde_json::Map::new(),
content: json!({"input": "finish"}),
};
let handle = llm_call(
LlmCallParams::builder()
.name("openai.responses")
.request(&request)
.attributes(LlmAttributes::STREAMING)
.build(),
)
.unwrap();
let mut wrapper = LlmStreamWrapper::new(
make_stream(vec![Ok(terminal_event.clone())]),
handle,
collector,
finalizer,
None,
None,
Some(Arc::new(OpenAIResponsesCodec)),
);

assert_eq!(wrapper.next().await.unwrap().unwrap(), terminal_event);
drop(wrapper);

let events = captured_snapshot(&events);
let end_event = events
.iter()
.find(|event| is_llm_end(event))
.expect("expected END event after the terminal response was dropped");
let metadata = end_event.metadata().unwrap();
assert_eq!(metadata["otel.status_code"], json!("OK"));
assert!(metadata.get("otel.status_description").is_none());

deregister_subscriber("stream_terminal_drop_status_test").unwrap();
}

#[tokio::test]
async fn stream_termination_modes_close_accounting_without_losing_evidence() {
let _lock = TEST_MUTEX.lock().unwrap();
Expand Down
21 changes: 18 additions & 3 deletions crates/core/tests/unit/stream_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,11 +30,26 @@ fn partial_stream_usage_is_not_treated_as_authoritative_without_terminal_evidenc
};
assert!(!has_authoritative_final_usage(Some(&partial)));

let terminal = AnnotatedLlmResponse {
finish_reason: Some(FinishReason::Complete),
for finish_reason in [
FinishReason::Complete,
FinishReason::Length,
FinishReason::ToolUse,
FinishReason::ContentFilter,
] {
let terminal = AnnotatedLlmResponse {
finish_reason: Some(finish_reason),
..partial.clone()
};
assert!(has_authoritative_final_usage(Some(&terminal)));
assert!(has_authoritative_terminal_outcome(Some(&terminal)));
}

let failed = AnnotatedLlmResponse {
finish_reason: Some(FinishReason::Unknown("failed".to_string())),
..partial
};
assert!(has_authoritative_final_usage(Some(&terminal)));
assert!(has_authoritative_final_usage(Some(&failed)));
assert!(!has_authoritative_terminal_outcome(Some(&failed)));
}

#[test]
Expand Down
Loading