feat: add rest sse streaming adapter

This commit is contained in:
a.tolmachev
2026-04-06 10:54:01 +03:00
parent 8204a59dac
commit bf56494336
10 changed files with 560 additions and 25 deletions
+177 -3
View File
@@ -7,7 +7,7 @@ use reqwest::{
};
use serde_json::Value;
use crate::{RestAdapterError, RestRequest, RestResponse};
use crate::{RestAdapterError, RestRequest, RestResponse, RestWindowRequest, RestWindowResponse};
#[derive(Clone, Debug)]
pub struct RestAdapter {
@@ -62,6 +62,61 @@ impl RestAdapter {
body,
})
}
pub async fn execute_window(
&self,
target: &RestTarget,
request: &RestWindowRequest,
) -> Result<RestWindowResponse, RestAdapterError> {
let url = build_url(target, &request.request)?;
let mut headers = build_headers(target, &request.request)?;
headers.insert(
reqwest::header::ACCEPT,
HeaderValue::from_static("text/event-stream"),
);
let mut builder = self
.client
.request(to_reqwest_method(target.method), url)
.headers(headers)
.timeout(Duration::from_millis(request.request.timeout_ms));
if let Some(body) = &request.request.body {
builder = builder.json(body);
}
let response = builder.send().await?;
let status = response.status();
if !status.is_success() {
let headers = normalize_headers(response.headers());
let body = decode_body(response).await?;
return Err(RestAdapterError::UnexpectedStatus {
status: status.as_u16(),
body: Value::Object(
[
(
"headers".to_owned(),
serde_json::to_value(headers).unwrap_or(Value::Null),
),
("body".to_owned(), body),
]
.into_iter()
.collect(),
),
});
}
let (status_code, headers, body) =
crate::sse::collect_sse_window(response, request.window_duration_ms, request.max_items)
.await?;
Ok(RestWindowResponse {
status_code,
headers,
body,
})
}
}
fn build_url(target: &RestTarget, request: &RestRequest) -> Result<reqwest::Url, RestAdapterError> {
@@ -172,13 +227,15 @@ mod tests {
Json, Router,
extract::{Path, Query},
http::HeaderMap,
response::sse::{Event, KeepAlive, Sse},
routing::{get, post},
};
use crank_core::{HttpMethod, RestTarget};
use futures_util::stream;
use serde_json::{Value, json};
use tokio::net::TcpListener;
use crate::{RestAdapter, RestAdapterError, RestRequest};
use crate::{RestAdapter, RestAdapterError, RestRequest, RestWindowRequest};
#[tokio::test]
async fn executes_rest_request_and_normalizes_json_response() {
@@ -242,10 +299,104 @@ mod tests {
));
}
#[tokio::test]
async fn collects_sse_events_with_window_bounds() {
let base_url = spawn_test_server().await;
let adapter = RestAdapter::new();
let target = RestTarget {
base_url,
method: HttpMethod::Get,
path_template: "/events".to_owned(),
static_headers: BTreeMap::new(),
};
let request = RestWindowRequest {
request: RestRequest {
path_params: BTreeMap::new(),
query_params: BTreeMap::new(),
headers: BTreeMap::new(),
body: None,
timeout_ms: 1_000,
},
window_duration_ms: 1_000,
max_items: Some(2),
};
let response = adapter.execute_window(&target, &request).await.unwrap();
assert_eq!(response.status_code, 200);
assert_eq!(
response.body,
json!({
"items": [
{ "message": "one" },
{ "message": "two" }
],
"done": false
})
);
}
#[tokio::test]
async fn returns_timeout_window_when_no_events_arrive_before_deadline() {
let base_url = spawn_test_server().await;
let adapter = RestAdapter::new();
let target = RestTarget {
base_url,
method: HttpMethod::Get,
path_template: "/events-idle".to_owned(),
static_headers: BTreeMap::new(),
};
let request = RestWindowRequest {
request: RestRequest {
path_params: BTreeMap::new(),
query_params: BTreeMap::new(),
headers: BTreeMap::new(),
body: None,
timeout_ms: 1_000,
},
window_duration_ms: 50,
max_items: Some(10),
};
let response = adapter.execute_window(&target, &request).await.unwrap();
assert_eq!(response.body, json!({ "items": [], "done": true }));
}
#[tokio::test]
async fn rejects_malformed_sse_payloads() {
let base_url = spawn_test_server().await;
let adapter = RestAdapter::new();
let target = RestTarget {
base_url,
method: HttpMethod::Get,
path_template: "/events-broken".to_owned(),
static_headers: BTreeMap::new(),
};
let request = RestWindowRequest {
request: RestRequest {
path_params: BTreeMap::new(),
query_params: BTreeMap::new(),
headers: BTreeMap::new(),
body: None,
timeout_ms: 1_000,
},
window_duration_ms: 1_000,
max_items: Some(10),
};
let error = adapter.execute_window(&target, &request).await.unwrap_err();
assert!(matches!(error, RestAdapterError::InvalidSseEvent));
}
async fn spawn_test_server() -> String {
let app = Router::new()
.route("/users/{user_id}", post(create_user))
.route("/fail", get(fail));
.route("/fail", get(fail))
.route("/events", get(sse_events))
.route("/events-idle", get(sse_idle))
.route("/events-broken", get(sse_broken));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
@@ -286,4 +437,27 @@ mod tests {
Json(json!({ "error": "upstream failed" })),
)
}
async fn sse_events()
-> Sse<impl futures_util::Stream<Item = Result<Event, std::convert::Infallible>>> {
let events = vec![
Ok(Event::default().data("{\"message\":\"one\"}")),
Ok(Event::default().data("{\"message\":\"two\"}")),
Ok(Event::default().data("{\"message\":\"three\"}")),
];
Sse::new(stream::iter(events)).keep_alive(KeepAlive::default())
}
async fn sse_idle()
-> Sse<impl futures_util::Stream<Item = Result<Event, std::convert::Infallible>>> {
Sse::new(stream::pending()).keep_alive(KeepAlive::default())
}
async fn sse_broken()
-> Sse<impl futures_util::Stream<Item = Result<Event, std::convert::Infallible>>> {
let events = vec![Ok(Event::default().data("{broken-json}"))];
Sse::new(stream::iter(events)).keep_alive(KeepAlive::default())
}
}
+4
View File
@@ -15,6 +15,10 @@ pub enum RestAdapterError {
InvalidHeaderValue { header: String },
#[error("request failed")]
Transport(#[from] reqwest::Error),
#[error("sse collection window expired before stream completed")]
WindowExpired,
#[error("rest endpoint returned status {status}")]
UnexpectedStatus { status: u16, body: Value },
#[error("sse stream produced malformed event payload")]
InvalidSseEvent,
}
+2 -1
View File
@@ -1,7 +1,8 @@
mod client;
mod error;
mod model;
mod sse;
pub use client::RestAdapter;
pub use error::RestAdapterError;
pub use model::{RestRequest, RestResponse};
pub use model::{RestRequest, RestResponse, RestWindowRequest, RestWindowResponse};
+17
View File
@@ -23,3 +23,20 @@ pub struct RestResponse {
pub headers: BTreeMap<String, String>,
pub body: Value,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Default)]
pub struct RestWindowRequest {
#[serde(flatten)]
pub request: RestRequest,
pub window_duration_ms: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_items: Option<u32>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct RestWindowResponse {
pub status_code: u16,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub headers: BTreeMap<String, String>,
pub body: Value,
}
+160
View File
@@ -0,0 +1,160 @@
use std::collections::BTreeMap;
use futures_util::StreamExt;
use reqwest::header::HeaderMap;
use serde_json::{Value, json};
use tokio::time::{Duration, Instant, timeout_at};
use crate::RestAdapterError;
pub async fn collect_sse_window(
response: reqwest::Response,
window_duration_ms: u64,
max_items: Option<u32>,
) -> Result<(u16, BTreeMap<String, String>, Value), RestAdapterError> {
let status = response.status();
let headers = normalize_headers(response.headers());
let deadline = Instant::now() + Duration::from_millis(window_duration_ms);
let mut stream = response.bytes_stream();
let mut buffer = String::new();
let mut items = Vec::new();
let mut done = true;
loop {
if max_items.is_some_and(|limit| items.len() >= limit as usize) {
done = false;
break;
}
let next_chunk = match timeout_at(deadline, stream.next()).await {
Ok(next_chunk) => next_chunk,
Err(_) => break,
};
let Some(next_chunk) = next_chunk else {
break;
};
let chunk = next_chunk?;
buffer.push_str(&String::from_utf8_lossy(&chunk));
while let Some(event_end) = find_event_boundary(&buffer) {
let event = buffer[..event_end].to_owned();
let boundary_len = boundary_length(&buffer[event_end..]);
buffer = buffer[event_end + boundary_len..].to_owned();
if let Some(item) = parse_sse_event(&event)? {
items.push(item);
if max_items.is_some_and(|limit| items.len() >= limit as usize) {
done = false;
break;
}
}
}
if !done {
break;
}
}
Ok((
status.as_u16(),
headers,
json!({
"items": items,
"done": done
}),
))
}
fn parse_sse_event(raw: &str) -> Result<Option<Value>, RestAdapterError> {
let mut data_lines = Vec::new();
for line in raw.lines() {
if line.is_empty() || line.starts_with(':') {
continue;
}
if let Some(value) = line.strip_prefix("data:") {
data_lines.push(value.trim_start().to_owned());
}
}
if data_lines.is_empty() {
return Ok(None);
}
let payload = data_lines.join("\n");
if payload.is_empty() {
return Ok(None);
}
serde_json::from_str::<Value>(&payload)
.map(Some)
.or_else(|_| {
if payload.starts_with('{') || payload.starts_with('[') {
Err(RestAdapterError::InvalidSseEvent)
} else {
Ok(Some(Value::String(payload)))
}
})
}
fn find_event_boundary(buffer: &str) -> Option<usize> {
buffer
.find("\r\n\r\n")
.or_else(|| buffer.find("\n\n"))
.or_else(|| buffer.find("\r\r"))
}
fn boundary_length(boundary: &str) -> usize {
if boundary.starts_with("\r\n\r\n") {
4
} else {
2
}
}
fn normalize_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
headers
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.as_str().to_owned(), value.to_owned()))
})
.collect()
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::parse_sse_event;
use crate::RestAdapterError;
#[test]
fn parses_json_event_payload() {
let event = "event: message\ndata: {\"message\":\"ok\"}\n\n";
let parsed = parse_sse_event(event).unwrap();
assert_eq!(parsed, Some(json!({ "message": "ok" })));
}
#[test]
fn ignores_comment_only_events() {
let parsed = parse_sse_event(": keepalive\n\n").unwrap();
assert_eq!(parsed, None);
}
#[test]
fn rejects_malformed_json_like_payload() {
let error = parse_sse_event("data: {broken-json}\n\n").unwrap_err();
assert!(matches!(error, RestAdapterError::InvalidSseEvent));
}
}