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
13 changes: 5 additions & 8 deletions crates/adaptive/tests/integration/runtime_integration_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,9 @@ use nemo_relay::api::llm::{
};
use nemo_relay::api::runtime::NemoRelayContextState;
use nemo_relay::api::runtime::global_context;
use nemo_relay::api::runtime::{LlmExecutionNextFn, LlmStreamExecutionNextFn, ToolExecutionNextFn};
use nemo_relay::api::runtime::{
LlmExecutionNextFn, LlmJsonStream, LlmStreamExecutionNextFn, ToolExecutionNextFn,
};
use nemo_relay::api::subscriber::{deregister_subscriber, flush_subscribers, register_subscriber};
use nemo_relay::api::tool::tool_call_execute;
use nemo_relay::codec::request::{AnnotatedLlmRequest, Message, MessageContent};
Expand Down Expand Up @@ -775,9 +777,7 @@ impl Plugin for HeaderPlugin {
}
chunks.push(Ok(chunk));
}
let stream = Box::pin(tokio_stream::iter(chunks))
as Pin<Box<dyn tokio_stream::Stream<Item = FlowResult<Json>> + Send>>;
Ok(stream)
Ok(LlmJsonStream::new(tokio_stream::iter(chunks)))
})
}),
)?;
Expand Down Expand Up @@ -853,10 +853,7 @@ async fn test_top_level_plugin_registers_request_and_execution_intercepts() {
let llm_stream_func: LlmStreamExecutionNextFn = Arc::new(|_req: LlmRequest| {
Box::pin(async move {
let chunks = vec![Ok(json!({"streamed": true}))];
Ok(Box::pin(tokio_stream::iter(chunks))
as Pin<
Box<dyn tokio_stream::Stream<Item = FlowResult<Json>> + Send>,
>)
Ok(LlmJsonStream::new(tokio_stream::iter(chunks)))
})
});
let collected = Arc::new(StdMutex::new(Vec::new()));
Expand Down
17 changes: 7 additions & 10 deletions crates/adaptive/tests/unit/acg_component_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,7 @@ use crate::config::AcgComponentConfig;
use crate::storage::memory::InMemoryBackend;
use crate::storage::traits::StorageBackendDyn;
use nemo_relay::api::llm::LlmRequest;
use nemo_relay::api::runtime::LlmExecutionNextFn;
use nemo_relay::api::runtime::LlmStreamExecutionNextFn;
use nemo_relay::api::runtime::{LlmExecutionNextFn, LlmJsonStream, LlmStreamExecutionNextFn};
use nemo_relay::codec::request::{AnnotatedLlmRequest, Message, MessageContent};
use serde_json::{Value, json};
use tokio_stream::StreamExt;
Expand Down Expand Up @@ -646,10 +645,9 @@ async fn acg_component_stream_execution_intercept_rewrites_streaming_requests()
);
let next: LlmStreamExecutionNextFn = Arc::new(|req| {
Box::pin(async move {
Ok(Box::pin(tokio_stream::iter(vec![Ok(req.content)]))
as Pin<
Box<dyn tokio_stream::Stream<Item = nemo_relay::error::Result<Json>> + Send>,
>)
Ok(LlmJsonStream::new(tokio_stream::iter(vec![
Ok(req.content),
])))
})
});

Expand Down Expand Up @@ -1295,10 +1293,9 @@ async fn acg_component_stream_execution_intercept_passes_original_request_when_t
);
let next: LlmStreamExecutionNextFn = Arc::new(|req| {
Box::pin(async move {
Ok(Box::pin(tokio_stream::iter(vec![Ok(req.content)]))
as Pin<
Box<dyn tokio_stream::Stream<Item = nemo_relay::error::Result<Json>> + Send>,
>)
Ok(LlmJsonStream::new(tokio_stream::iter(vec![
Ok(req.content),
])))
})
});

Expand Down
15 changes: 5 additions & 10 deletions crates/adaptive/tests/unit/runtime_features_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ use nemo_relay::api::registry::{
register_llm_execution_intercept, register_llm_request_intercept,
register_llm_stream_execution_intercept, register_tool_execution_intercept,
};
use nemo_relay::api::runtime::LlmJsonStream;
use nemo_relay::api::runtime::ToolExecutionNextFn;
use nemo_relay::api::runtime::global_context;
use nemo_relay::api::runtime::{
Expand Down Expand Up @@ -746,14 +747,9 @@ async fn registration_context_registers_all_supported_callback_types() {
7,
Arc::new(|_name, request, _next| {
Box::pin(async move {
Ok(Box::pin(tokio_stream::iter(vec![Ok(request.content)]))
as Pin<
Box<
dyn tokio_stream::Stream<
Item = nemo_relay::error::Result<nemo_relay::json::Json>,
> + Send,
>,
>)
Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok(
request.content
)])))
})
}),
)
Expand Down Expand Up @@ -878,8 +874,7 @@ async fn acg_feature_registers_execution_and_stream_intercepts() {

let stream_next: LlmStreamExecutionNextFn = Arc::new(|request| {
Box::pin(async move {
let stream: nemo_relay::api::runtime::LlmJsonStream =
Box::pin(tokio_stream::iter(vec![Ok(request.content)]));
let stream = LlmJsonStream::new(tokio_stream::iter(vec![Ok(request.content)]));
Ok(stream)
})
});
Expand Down
2 changes: 1 addition & 1 deletion crates/cli/src/gateway/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -534,7 +534,7 @@ fn sse_json_stream(response: reqwest::Response) -> LlmJsonStream {
Err(error) => yield Err(error),
}
};
Box::pin(stream)
LlmJsonStream::new(stream)
}

// Re-encodes a runtime JSON stream as `text/event-stream` frames for the downstream client. Event
Expand Down
6 changes: 3 additions & 3 deletions crates/cli/tests/coverage/shared/gateway_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1776,7 +1776,7 @@ async fn streaming_gateway_call_guard_finishes_when_body_is_dropped() {
.await
.unwrap();

let stream: LlmJsonStream = Box::pin(futures_util::stream::pending::<
let stream = LlmJsonStream::new(futures_util::stream::pending::<
std::result::Result<Value, FlowError>,
>());
let body = client_sse_body(
Expand Down Expand Up @@ -1856,7 +1856,7 @@ fn streaming_gateway_call_guard_finishes_without_a_current_runtime() {
(manager, prep)
});
let final_response = json!({ "output_text": "streamed final" });
let stream: LlmJsonStream = Box::pin(futures_util::stream::pending::<
let stream = LlmJsonStream::new(futures_util::stream::pending::<
std::result::Result<Value, FlowError>,
>());
let body = client_sse_body(
Expand Down Expand Up @@ -1934,7 +1934,7 @@ async fn streaming_body_records_final_response_for_turn_output() {
let session_id = prep.session_id.clone();
let owner_subagent_id = prep.owner_subagent_id.clone();
let final_response = json!({ "output_text": "streamed final" });
let stream: LlmJsonStream = Box::pin(futures_util::stream::empty::<
let stream = LlmJsonStream::new(futures_util::stream::empty::<
std::result::Result<Value, FlowError>,
>());
let body = client_sse_body(
Expand Down
2 changes: 1 addition & 1 deletion crates/core/src/api/llm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1173,7 +1173,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu
response_codec,
lifecycle_subscribers,
);
Ok(Box::pin(wrapper) as LlmJsonStream)
Ok(LlmJsonStream::from_closeable(wrapper))
}
Err(error) => {
let end_metadata =
Expand Down
4 changes: 2 additions & 2 deletions crates/core/src/api/runtime.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,8 @@ pub mod subscriber_dispatcher;
pub use callbacks::{
EventSanitizeFn, EventSubscriberFn, LlmCollectorFn, LlmConditionalFn, LlmExecutionFn,
LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestFn,
LlmSanitizeResponseFn, LlmStreamExecutionFn, LlmStreamExecutionNextFn, ToolConditionalFn,
ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn,
LlmSanitizeResponseFn, LlmStreamExecutionFn, LlmStreamExecutionNextFn, LlmStreamInner,
ToolConditionalFn, ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn,
};
pub use global::global_context;
pub use scope_stack::{
Expand Down
81 changes: 80 additions & 1 deletion crates/core/src/api/runtime/callbacks.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};

use tokio_stream::Stream;

Expand Down Expand Up @@ -240,7 +241,85 @@ pub type LlmExecutionFn = Arc<
+ Sync,
>;
/// Stream of JSON chunks produced by the managed streaming LLM pipeline.
pub type LlmJsonStream = Pin<Box<dyn Stream<Item = Result<Json>> + Send>>;
///
/// In addition to ordinary stream polling, managed streams provide an explicit
/// asynchronous close operation. A successful close means the producer has
/// released its resources; subsequent polls return no more chunks.
pub struct LlmJsonStream {
inner: Pin<Box<dyn LlmStreamInner>>,
}

impl LlmJsonStream {
/// Wrap a stream whose producer has no asynchronous teardown work.
pub fn new<S>(stream: S) -> Self
where
S: Stream<Item = Result<Json>> + Send + 'static,
{
Self {
inner: Box::pin(DefaultLlmStream {
stream: Some(Box::pin(stream)),
}),
}
}

/// Wrap a stream that implements explicit asynchronous teardown.
pub fn from_closeable<S>(stream: S) -> Self
where
S: LlmStreamInner + 'static,
{
Self {
inner: Box::pin(stream),
}
}

/// Stop the producer and wait for its cleanup to complete.
pub async fn close(&mut self) -> Result<()> {
self.inner.as_mut().close().await
}
}

impl Stream for LlmJsonStream {
type Item = Result<Json>;

fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.get_mut().inner.as_mut().poll_next(cx)
}
}

/// Internal close-aware stream implementation.
pub trait LlmStreamInner: Stream<Item = Result<Json>> + Send {
/// Stop the producer and wait for cleanup. Implementations must be idempotent.
fn close(self: Pin<&mut Self>) -> Pin<Box<dyn Future<Output = Result<()>> + Send + '_>>;
}

struct DefaultLlmStream<S> {
stream: Option<Pin<Box<S>>>,
}

impl<S> Stream for DefaultLlmStream<S>
where
S: Stream<Item = Result<Json>> + Send,
{
type Item = Result<Json>;

fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
match this.stream.as_mut() {
Some(stream) => stream.as_mut().poll_next(cx),
None => Poll::Ready(None),
}
}
}

impl<S> LlmStreamInner for DefaultLlmStream<S>
where
S: Stream<Item = Result<Json>> + Send,
{
fn close(self: Pin<&mut Self>) -> Pin<Box<dyn Future<Output = Result<()>> + Send + '_>> {
self.get_mut().stream.take();
Box::pin(async { Ok(()) })
}
}
/// Per-chunk collector used by the streaming LLM runtime.
///
/// # Parameters
Expand Down
2 changes: 1 addition & 1 deletion crates/core/src/error.rs
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ impl std::fmt::Display for UpstreamFailure {
///
/// Each variant represents a distinct failure mode that callers can match on
/// to determine the appropriate recovery strategy.
#[derive(Debug, Error)]
#[derive(Clone, Debug, Error)]
pub enum FlowError {
/// A resource with the given name is already registered.
///
Expand Down
4 changes: 2 additions & 2 deletions crates/core/src/plugin/dynamic/native.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2159,11 +2159,11 @@ fn native_stream_to_relay_stream(
next_ctx: Option<NativeStreamNextContext>,
callback_user_data: Option<Arc<NativeCallbackUserData>>,
) -> FlowResult<LlmJsonStream> {
Ok(Box::pin(NativeRelayLlmStream::from_raw(
Ok(LlmJsonStream::new(NativeRelayLlmStream::from_raw(
raw,
next_ctx,
callback_user_data,
)?) as LlmJsonStream)
)?))
}

fn drop_native_stream(mut raw: NemoRelayNativeLlmStreamV1) {
Expand Down
4 changes: 3 additions & 1 deletion crates/core/src/plugin/dynamic/worker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1745,7 +1745,9 @@ impl WorkerPluginCallback {
}
guard.finish();
});
Ok(Box::pin(tokio_stream::wrappers::ReceiverStream::new(rx)))
Ok(LlmJsonStream::new(
tokio_stream::wrappers::ReceiverStream::new(rx),
))
}

fn base_request(
Expand Down
Loading
Loading