aether-llm 0.9.13

Multi-provider LLM abstraction layer for the Aether AI agent framework
Documentation
use crate::providers::http::HttpResponseMetadata;
use crate::{LlmError, LlmResponse, LlmResponseStream, ProviderError, Result, StopReason, ToolCallRequest};
use async_stream::{stream, try_stream};
use futures::{Stream, StreamExt};
use std::collections::BTreeMap;
use std::fmt::Debug;
use std::future::Future;
use std::pin::pin;
use std::time::Duration;
use tracing::warn;

pub(crate) fn response_stream<T, U, V>(
    open: impl Future<Output = Result<OpenedStream<T>>> + Send + 'static,
    decode: impl FnMut(U, &mut StreamAssembler<V>) -> Result<Vec<LlmResponse>> + Send + 'static,
    idle_timeout: Duration,
) -> LlmResponseStream
where
    T: Stream<Item = Result<U>> + Send + 'static,
    U: Send,
    V: Ord + Debug + Send + 'static,
{
    Box::pin(stream! {
        match within(idle_timeout, open).await.flatten() {
            Ok(OpenedStream { events, http }) => {
                let mut responses = pin!(decode_events(events, decode, idle_timeout));
                while let Some(response) = responses.next().await {
                    yield response.map_err(|error| match &http {
                        Some(http) => http.annotate(error),
                        None => error,
                    });
                }
            }
            Err(error) => yield Err(error),
        }
    })
}

pub(crate) fn error_stream(error: LlmError) -> LlmResponseStream {
    Box::pin(futures::stream::iter([Err(error)]))
}

pub(crate) struct OpenedStream<T> {
    events: T,
    http: Option<HttpResponseMetadata>,
}

impl<T> OpenedStream<T> {
    pub fn new(events: T) -> Self {
        Self { events, http: None }
    }

    pub fn with_http(events: T, http: HttpResponseMetadata) -> Self {
        Self { events, http: Some(http) }
    }
}

pub(crate) struct StreamAssembler<T> {
    tool_calls: BTreeMap<T, ToolCallRequest>,
    stop_reason: Option<StopReason>,
    termination: Termination,
}

impl<T: Ord + Debug> StreamAssembler<T> {
    pub fn new() -> Self {
        Self { tool_calls: BTreeMap::new(), stop_reason: None, termination: Termination::Pending }
    }

    pub fn start_tool(&mut self, index: T, id: String, name: String) -> LlmResponse {
        let start = LlmResponse::ToolRequestStart { id: id.clone(), name: name.clone() };
        self.tool_calls.insert(index, ToolCallRequest { id, name, arguments: String::new() });
        start
    }

    pub fn append_tool_args(&mut self, index: &T, chunk: String) -> Option<LlmResponse> {
        if chunk.is_empty() {
            return None;
        }

        let Some(tool_call) = self.tool_calls.get_mut(index) else {
            warn!("Received tool call arguments for unknown index {index:?}");
            return None;
        };
        tool_call.arguments.push_str(&chunk);
        Some(LlmResponse::ToolRequestArg { id: tool_call.id.clone(), chunk })
    }

    pub fn complete_tool(&mut self, index: &T) -> Option<LlmResponse> {
        self.tool_calls.remove(index).map(|tool_call| LlmResponse::ToolRequestComplete { tool_call })
    }

    pub fn complete_tool_with(&mut self, index: &T, mut tool_call: ToolCallRequest) -> Vec<LlmResponse> {
        let mut responses = Vec::new();
        match self.tool_calls.remove(index) {
            Some(streamed) if tool_call.arguments.is_empty() => tool_call.arguments = streamed.arguments,
            Some(_) => {}
            None => {
                responses
                    .push(LlmResponse::ToolRequestStart { id: tool_call.id.clone(), name: tool_call.name.clone() });
            }
        }
        responses.push(LlmResponse::ToolRequestComplete { tool_call });
        responses
    }

    pub fn complete_all_tools(&mut self) -> Vec<LlmResponse> {
        std::mem::take(&mut self.tool_calls)
            .into_values()
            .map(|tool_call| LlmResponse::ToolRequestComplete { tool_call })
            .collect()
    }

    pub fn stop(&mut self, stop_reason: StopReason) {
        self.stop_reason = Some(stop_reason);
    }

    pub fn allow_eof(&mut self) {
        self.termination = Termination::AtEof;
    }

    pub fn finish_now(&mut self) {
        self.termination = Termination::Immediate;
    }

    fn finish(mut self) -> Result<Vec<LlmResponse>> {
        if matches!(self.termination, Termination::Pending) {
            return Err(ProviderError::stream_interrupted("stream ended before the provider's terminal event").into());
        }

        let mut responses = self.complete_all_tools();
        responses.push(LlmResponse::Done { stop_reason: self.stop_reason });
        Ok(responses)
    }
}

enum Termination {
    Pending,
    AtEof,
    Immediate,
}

fn decode_events<T, U>(
    events: impl Stream<Item = Result<T>> + Send,
    mut decode: impl FnMut(T, &mut StreamAssembler<U>) -> Result<Vec<LlmResponse>> + Send,
    idle_timeout: Duration,
) -> impl Stream<Item = Result<LlmResponse>> + Send
where
    T: Send,
    U: Ord + Debug + Send,
{
    try_stream! {
        let mut events = pin!(events);
        let mut turn = StreamAssembler::new();
        yield LlmResponse::Start;

        while let Some(event) = within(idle_timeout, events.next()).await? {
            for response in decode(event?, &mut turn)? {
                yield response;
            }

            if matches!(turn.termination, Termination::Immediate) {
                break;
            }
        }

        for response in turn.finish()? {
            yield response;
        }
    }
}

async fn within<T>(idle_timeout: Duration, future: impl Future<Output = T>) -> Result<T> {
    tokio::time::timeout(idle_timeout, future)
        .await
        .map_err(|_| ProviderError::timeout(format!("Provider sent nothing for {}s", idle_timeout.as_secs())).into())
}

#[cfg(test)]
mod tests {
    use super::*;
    use futures::stream;
    use std::future::{pending, ready};

    #[tokio::test(start_paused = true)]
    async fn provider_that_never_responds_times_out() {
        let responses = collect(text_stream(pending::<Result<OpenedStream<stream::Empty<_>>>>())).await;

        assert_eq!(responses.len(), 1, "{responses:?}");
        assert_timed_out(&responses);
    }

    #[tokio::test(start_paused = true)]
    async fn provider_that_goes_silent_mid_stream_times_out() {
        let events = stream::iter([Ok("partial")]).chain(stream::pending());

        let responses = collect(text_stream(ready(Ok(OpenedStream::new(events))))).await;

        assert!(matches!(&responses[..2], [Ok(LlmResponse::Start), Ok(LlmResponse::Text { .. })]), "{responses:?}");
        assert_timed_out(&responses);
    }

    #[tokio::test(start_paused = true)]
    async fn slow_stream_that_keeps_sending_events_completes() {
        let events = stream::iter(["one", "two", "three"]).then(|text| async move {
            tokio::time::sleep(Duration::from_secs(50)).await;
            Ok(text)
        });

        let responses = collect(text_stream(ready(Ok(OpenedStream::new(events))))).await;

        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
        assert!(matches!(responses.last(), Some(Ok(LlmResponse::Done { .. }))), "{responses:?}");
    }

    #[tokio::test]
    async fn errors_after_opening_carry_the_http_diagnostics() {
        let events = stream::iter([Err(ProviderError::stream_interrupted("connection reset").into())]);
        let http = HttpResponseMetadata { status: 200, request_id: Some("req-1".to_string()) };

        let responses = collect(text_stream(ready(Ok(OpenedStream::with_http(events, http))))).await;

        let error = responses.last().and_then(|response| response.as_ref().err()).and_then(LlmError::provider);
        assert_eq!(error.and_then(|error| error.request_id.as_deref()), Some("req-1"), "{responses:?}");
        assert_eq!(error.and_then(|error| error.http_status), Some(200));
    }

    const IDLE_TIMEOUT: Duration = Duration::from_mins(1);

    fn text_stream<T>(open: impl Future<Output = Result<OpenedStream<T>>> + Send + 'static) -> LlmResponseStream
    where
        T: Stream<Item = Result<&'static str>> + Send + 'static,
    {
        let decode = |text: &'static str, turn: &mut StreamAssembler<u32>| {
            turn.allow_eof();
            Ok(vec![LlmResponse::Text { chunk: text.to_string() }])
        };
        response_stream(open, decode, IDLE_TIMEOUT)
    }

    async fn collect(responses: LlmResponseStream) -> Vec<Result<LlmResponse>> {
        responses.collect().await
    }

    fn assert_timed_out(responses: &[Result<LlmResponse>]) {
        let kind = responses.last().and_then(|response| response.as_ref().err()).and_then(LlmError::provider);
        assert_eq!(kind.map(|error| error.kind), Some(crate::ProviderErrorKind::Timeout), "{responses:?}");
        assert!(!responses.iter().any(|response| matches!(response, Ok(LlmResponse::Done { .. }))), "{responses:?}");
    }
}