use std::collections::VecDeque;
use std::pin::Pin;
use std::task::{Context, Poll};
use futures_util::{Stream, StreamExt};
use reqwest::header::{ACCEPT, AUTHORIZATION, CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue};
use crate::ag_ui::consumer::{ProtocolError, RunConsumer, RunResult};
use crate::ag_ui::{Event, RunAgentInput};
pub const DEFAULT_MAX_EVENT_BYTES: usize = 4 * 1024 * 1024;
#[derive(Debug)]
pub enum ClientError {
Http(reqwest::Error),
Status {
status: u16,
body: String,
},
ContentType(String),
InvalidHeader(String),
EventTooLarge { limit: usize },
InvalidJson(serde_json::Error),
Protocol(ProtocolError),
}
impl std::fmt::Display for ClientError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Http(err) => write!(f, "AG-UI request failed: {err}"),
Self::Status { status, body } => {
write!(f, "AG-UI agent answered HTTP {status}: {body}")
}
Self::ContentType(found) => write!(
f,
"AG-UI agent answered with content type '{found}', expected text/event-stream"
),
Self::InvalidHeader(name) => write!(f, "invalid AG-UI request header '{name}'"),
Self::EventTooLarge { limit } => {
write!(f, "AG-UI event exceeded the {limit}-byte limit")
}
Self::InvalidJson(err) => write!(f, "AG-UI event is not valid JSON: {err}"),
Self::Protocol(err) => err.fmt(f),
}
}
}
impl std::error::Error for ClientError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Http(err) => Some(err),
Self::InvalidJson(err) => Some(err),
Self::Protocol(err) => Some(err),
_ => None,
}
}
}
impl From<ProtocolError> for ClientError {
fn from(err: ProtocolError) -> Self {
Self::Protocol(err)
}
}
#[derive(Clone, Debug)]
pub struct AgUiClient {
http: reqwest::Client,
url: String,
bearer_token: Option<String>,
headers: HeaderMap,
max_event_bytes: usize,
}
impl AgUiClient {
pub fn new(url: impl Into<String>) -> Self {
Self {
http: reqwest::Client::new(),
url: url.into(),
bearer_token: None,
headers: HeaderMap::new(),
max_event_bytes: DEFAULT_MAX_EVENT_BYTES,
}
}
pub fn with_http_client(mut self, http: reqwest::Client) -> Self {
self.http = http;
self
}
pub fn with_bearer_token(mut self, token: impl Into<String>) -> Self {
self.bearer_token = Some(token.into());
self
}
pub fn with_header(mut self, name: &str, value: &str) -> Result<Self, ClientError> {
let header = HeaderName::from_bytes(name.as_bytes())
.map_err(|_| ClientError::InvalidHeader(name.to_owned()))?;
let mut value = HeaderValue::from_str(value)
.map_err(|_| ClientError::InvalidHeader(name.to_owned()))?;
value.set_sensitive(true);
self.headers.insert(header, value);
Ok(self)
}
pub fn with_max_event_bytes(mut self, limit: usize) -> Self {
self.max_event_bytes = limit;
self
}
pub fn url(&self) -> &str {
&self.url
}
pub async fn run(&self, input: &RunAgentInput) -> Result<EventStream, ClientError> {
let mut request = self
.http
.post(&self.url)
.headers(self.headers.clone())
.header(ACCEPT, "text/event-stream")
.json(input);
if let Some(token) = &self.bearer_token {
let mut value = HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(|_| ClientError::InvalidHeader(AUTHORIZATION.to_string()))?;
value.set_sensitive(true);
request = request.header(AUTHORIZATION, value);
}
let response = request.send().await.map_err(ClientError::Http)?;
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
return Err(ClientError::Status {
status: status.as_u16(),
body: body.chars().take(2048).collect(),
});
}
let content_type = response
.headers()
.get(CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or_default()
.to_owned();
if !content_type
.to_ascii_lowercase()
.starts_with("text/event-stream")
{
return Err(ClientError::ContentType(content_type));
}
let body = response
.bytes_stream()
.map(|chunk| chunk.map(|bytes| bytes.to_vec()));
Ok(EventStream {
body: Box::pin(body),
sse: SseReader::new(self.max_event_bytes),
consumer: Some(RunConsumer::for_input(input)),
finished: None,
queue: VecDeque::new(),
done: false,
})
}
}
type ByteStream = Pin<Box<dyn Stream<Item = Result<Vec<u8>, reqwest::Error>> + Send>>;
pub struct EventStream {
body: ByteStream,
sse: SseReader,
consumer: Option<RunConsumer>,
finished: Option<RunResult>,
queue: VecDeque<Event>,
done: bool,
}
impl std::fmt::Debug for EventStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EventStream")
.field("queued", &self.queue.len())
.field("done", &self.done)
.finish_non_exhaustive()
}
}
impl EventStream {
pub fn result(&self) -> Option<&RunResult> {
self.finished
.as_ref()
.or_else(|| self.consumer.as_ref().map(RunConsumer::result))
}
pub async fn into_result(mut self) -> Result<RunResult, ClientError> {
while let Some(event) = self.next().await {
event?;
}
self.finished.take().ok_or_else(|| {
ClientError::Protocol(ProtocolError::new("the stream ended without a result"))
})
}
fn fail(&mut self, err: ClientError) -> Poll<Option<Result<Event, ClientError>>> {
self.done = true;
self.consumer = None;
self.queue.clear();
Poll::Ready(Some(Err(err)))
}
fn accept(&mut self, data: &str) -> Result<(), ClientError> {
let value: serde_json::Value =
serde_json::from_str(data).map_err(ClientError::InvalidJson)?;
if let Some(consumer) = self.consumer.as_mut() {
self.queue.extend(consumer.push_value(value)?);
}
Ok(())
}
fn end(&mut self) -> Result<(), ClientError> {
self.sse.flush()?;
while let Some(data) = self.sse.ready.pop_front() {
self.accept(&data)?;
}
if let Some(consumer) = self.consumer.take() {
self.finished = Some(consumer.finish()?);
}
Ok(())
}
}
impl Stream for EventStream {
type Item = Result<Event, ClientError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
loop {
if let Some(event) = this.queue.pop_front() {
return Poll::Ready(Some(Ok(event)));
}
if this.done {
return Poll::Ready(None);
}
if let Some(data) = this.sse.ready.pop_front() {
if let Err(err) = this.accept(&data) {
return this.fail(err);
}
continue;
}
match this.body.as_mut().poll_next(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Some(Ok(chunk))) => {
if let Err(err) = this.sse.feed(&chunk) {
return this.fail(err);
}
}
Poll::Ready(Some(Err(err))) => return this.fail(ClientError::Http(err)),
Poll::Ready(None) => {
if let Err(err) = this.end() {
return this.fail(err);
}
this.done = true;
}
}
}
}
}
#[derive(Debug)]
struct SseReader {
buffer: Vec<u8>,
data: Option<String>,
ready: VecDeque<String>,
limit: usize,
}
impl SseReader {
fn new(limit: usize) -> Self {
Self {
buffer: Vec::new(),
data: None,
ready: VecDeque::new(),
limit,
}
}
fn feed(&mut self, chunk: &[u8]) -> Result<(), ClientError> {
self.buffer.extend_from_slice(chunk);
let mut start = 0;
while let Some(offset) = self.buffer[start..]
.iter()
.position(|b| *b == b'\n' || *b == b'\r')
{
let end = start + offset;
let next = match (self.buffer[end], self.buffer.get(end + 1)) {
(b'\r', Some(b'\n')) => end + 2,
(b'\r', None) => break,
_ => end + 1,
};
let line = String::from_utf8_lossy(&self.buffer[start..end]).into_owned();
start = next;
self.line(&line)?;
}
self.buffer.drain(..start);
let pending = self.buffer.len() + self.data.as_ref().map_or(0, String::len);
if pending > self.limit {
return Err(ClientError::EventTooLarge { limit: self.limit });
}
Ok(())
}
fn line(&mut self, line: &str) -> Result<(), ClientError> {
if line.is_empty() {
if let Some(data) = self.data.take() {
self.ready.push_back(data);
}
return Ok(());
}
let (field, value) = match line.split_once(':') {
Some((field, value)) => (field, value.strip_prefix(' ').unwrap_or(value)),
None => (line, ""),
};
if field == "data" {
match &mut self.data {
Some(data) => {
data.push('\n');
data.push_str(value);
}
None => self.data = Some(value.to_owned()),
}
if self.data.as_ref().map_or(0, String::len) > self.limit {
return Err(ClientError::EventTooLarge { limit: self.limit });
}
}
Ok(())
}
fn flush(&mut self) -> Result<(), ClientError> {
let rest = std::mem::take(&mut self.buffer);
let rest = String::from_utf8_lossy(&rest);
let rest = rest.strip_suffix('\r').unwrap_or(&rest);
if !rest.is_empty() {
self.line(rest)?;
}
if let Some(data) = self.data.take() {
self.ready.push_back(data);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used, clippy::expect_used)]
use super::*;
fn read(chunks: &[&[u8]]) -> Vec<String> {
let mut reader = SseReader::new(1024);
for chunk in chunks {
reader.feed(chunk).unwrap();
}
reader.flush().unwrap();
reader.ready.drain(..).collect()
}
#[test]
fn sse_reader_handles_line_endings_and_split_frames() {
assert_eq!(read(&[b"data: a\n\ndata: b\r\n\r\n"]), ["a", "b"]);
assert_eq!(read(&[b"data: a\r", b"\n\r\n"]), ["a"]);
assert_eq!(read(&[b"data: a\r\rdata: b\r\r"]), ["a", "b"]);
assert_eq!(read(&[b"da", b"ta: {\"x\"", b":1}\n\n"]), [r#"{"x":1}"#]);
assert_eq!(
read(&[b": comment\nevent: x\ndata: 1\ndata: 2\n\n"]),
["1\n2"]
);
assert_eq!(read(&[b"data: tail"]), ["tail"]);
}
#[test]
fn sse_reader_bounds_one_event() {
let mut reader = SseReader::new(8);
assert!(matches!(
reader.feed(b"data: 0123456789\n"),
Err(ClientError::EventTooLarge { limit: 8 })
));
let mut reader = SseReader::new(8);
assert!(reader.feed(b"data: 0123456789").is_err());
}
}