commit 3f6a0fd79c5c9f294536810a2edb2f5419acf326
parent b0ed224ea47e84849faaa10ee7f971b1d554f1df
Author: Antoine A <>
Date: Tue, 29 Sep 2026 20:51:30 +0200
common: fix db long-pool, use sse-core crate for sse protocol and improve header parsing helper
Diffstat:
12 files changed, 100 insertions(+), 231 deletions(-)
diff --git a/adapters/taler-cyclos/src/cyclos_api/api.rs b/adapters/taler-cyclos/src/cyclos_api/api.rs
@@ -19,6 +19,7 @@ use std::borrow::Cow;
use http_client::{
ApiErr, Client, ClientErr, Ctx,
builder::{Req, Res},
+ headers::HeaderParser as _,
sse::SseClient,
};
use hyper::{
@@ -116,7 +117,7 @@ impl<'a> CyclosRequest<'a> {
}
}
- pub async fn into_sse(mut self, client: &mut SseClient) -> ApiResult<()> {
+ pub async fn into_sse(mut self, client: &mut SseClient) -> ApiResult<bool> {
self.req = self.req.req_sse(client);
let (ctx, res) = self.send().await?;
res.sse(client).map_err(|e| ctx.wrap(e.into()))
@@ -137,8 +138,8 @@ impl<'a> CyclosRequest<'a> {
let (ctx, res) = self.send().await?;
async {
let res = Self::error_handling(res).await?;
- let current_page = res.int_header("x-current-page")?;
- let has_next_page = res.bool_header("x-has-next-page")?;
+ let current_page = res.int_header(HeaderName::from_static("x-current-page"))?;
+ let has_next_page = res.bool_header(HeaderName::from_static("x-has-next-page"))?;
Ok(Pagination {
page: res.json().await?,
current_page,
diff --git a/adapters/taler-cyclos/src/cyclos_api/client.rs b/adapters/taler-cyclos/src/cyclos_api/client.rs
@@ -126,7 +126,7 @@ impl Client<'_> {
&self,
client_id: i64,
sse_client: &mut SseClient,
- ) -> ApiResult<()> {
+ ) -> ApiResult<bool> {
self.request(Method::GET, "push/subscribe")
.query("clientId", client_id)
.query("kinds", "newNotification")
diff --git a/adapters/taler-cyclos/src/notification.rs b/adapters/taler-cyclos/src/notification.rs
@@ -25,7 +25,7 @@ use crate::cyclos_api::{
types::{NotificationEntityType, NotificationStatus},
};
-pub async fn watch_notification(client: &Client<'_>, notify: &Notify) -> ! {
+pub async fn watch_notification(client: &Client<'_>, notify: &Notify) -> () {
let client_id = Timestamp::now().as_microsecond();
let mut sse_client = SseClient::new();
let mut jitter = ExpoBackoffDecorr::default();
@@ -33,9 +33,12 @@ pub async fn watch_notification(client: &Client<'_>, notify: &Notify) -> ! {
let res: anyhow::Result<()> = async {
loop {
// Register listener
- client
+ let connected = client
.push_notifications(client_id, &mut sse_client)
.await?;
+ if !connected {
+ return Ok(())
+ }
jitter.reset();
// Read available ones
while let Some(message) = sse_client.next().await {
@@ -57,8 +60,12 @@ pub async fn watch_notification(client: &Client<'_>, notify: &Notify) -> ! {
}
}
.await;
- let err = res.unwrap_err();
- error!(target: "notification", "{err}");
- tokio::time::sleep(jitter.backoff()).await;
+ if let Err(err) = res {
+ error!(target: "notification", "{err}");
+ tokio::time::sleep(jitter.backoff()).await;
+ } else {
+ error!(target: "notification", "SSE connection refused with HTTP 204");
+ std::future::pending::<()>().await;
+ }
}
}
diff --git a/adapters/taler-cyclos/src/worker.rs b/adapters/taler-cyclos/src/worker.rs
@@ -175,7 +175,6 @@ pub async fn run_worker(
result = worker => result?,
_ = watcher => unreachable!("notification listener does not return"),
}
- Ok(())
}
pub struct Worker<'a> {
diff --git a/adapters/taler-magnet-bank/src/magnet_api/api.rs b/adapters/taler-magnet-bank/src/magnet_api/api.rs
@@ -19,6 +19,7 @@ use std::borrow::Cow;
use http_client::{
ApiErr, Client, ClientErr, Ctx,
builder::{Req, Res},
+ headers::HeaderParser as _,
};
use hyper::{Method, StatusCode, header};
use serde::{Deserialize, Serialize, de::DeserializeOwned};
@@ -68,7 +69,7 @@ async fn error_handling(res: Res) -> Result<Res, MagnetErr> {
StatusCode::OK => Ok(res),
StatusCode::BAD_REQUEST => Err(MagnetErr::Status(status)),
StatusCode::FORBIDDEN => {
- let cause = res.str_header(header::WWW_AUTHENTICATE.as_str())?;
+ let cause = res.str_header(header::WWW_AUTHENTICATE)?;
Err(MagnetErr::StatusCause(status, cause))
}
_ => {
diff --git a/adapters/taler-wise/src/wise_api/client.rs b/adapters/taler-wise/src/wise_api/client.rs
@@ -16,7 +16,7 @@
use std::{borrow::Cow, str::FromStr as _, sync::LazyLock};
-use http_client::{ApiErr, ClientErr, builder::Req};
+use http_client::{ApiErr, ClientErr, builder::Req, headers::HeaderParser as _};
use hyper::{
Method, StatusCode,
header::{AUTHORIZATION, HeaderValue, RETRY_AFTER},
@@ -73,7 +73,7 @@ impl<'a> Client<'a> {
res.json().await.map_err(|e| ctx.wrap(e.into()))
}
StatusCode::TOO_MANY_REQUESTS => Err(ctx.wrap(WiseErr::TooManyRequest(
- res.str_header(RETRY_AFTER.as_str()).unwrap_or_default(),
+ res.str_header(RETRY_AFTER).unwrap_or_default(),
))),
other => Err(ctx.wrap(WiseErr::Status(other))),
}
diff --git a/common/http-client/Cargo.toml b/common/http-client/Cargo.toml
@@ -34,4 +34,5 @@ tokio-util = { version = "0.7.17", default-features = false, features = [
"io",
] }
futures-util = { version = "0.3", default-features = false }
+sse-core = "0.2"
diff --git a/common/http-client/src/builder.rs b/common/http-client/src/builder.rs
@@ -23,7 +23,7 @@ use http::{
HeaderMap, HeaderName, HeaderValue, StatusCode,
header::{self},
};
-use http_body_util::{BodyDataStream, BodyExt, Full, Limited};
+use http_body_util::{BodyExt, Full, Limited};
use hyper::{Method, body::Bytes};
use serde::{Serialize, de::DeserializeOwned};
use serde_path_to_error::Track;
@@ -179,7 +179,7 @@ impl Req {
HeaderValue::from_static("text/event-stream"),
)
.header(header::CACHE_CONTROL, HeaderValue::from_static("no-cache"));
- if let Some(id) = &client.last_event_id {
+ if let Some(id) = &client.last_event_id() {
self = self.header(
HeaderName::from_static("last-event-id"),
HeaderValue::from_str(id).unwrap(),
@@ -239,37 +239,25 @@ impl Res {
&self.head.headers
}
- pub fn str_header(&self, name: &'static str) -> Result<String, ClientErr> {
- self.head
- .headers
- .str_header(name)
- .map_err(ClientErr::Headers)
- }
-
- pub fn int_header(&self, name: &'static str) -> Result<u64, ClientErr> {
- self.head
- .headers
- .int_header(name)
- .map_err(ClientErr::Headers)
- }
-
- pub fn bool_header(&self, name: &'static str) -> Result<bool, ClientErr> {
- self.head
- .headers
- .bool_header(name)
- .map_err(ClientErr::Headers)
- }
-
- pub fn sse(self, client: &mut SseClient) -> Result<(), ClientErr> {
- // TODO check content type?
- // TODO check status
- client.connect(BodyDataStream::new(self.body));
- Ok(())
+ pub fn sse(self, client: &mut SseClient) -> Result<bool, ClientErr> {
+ match self.status() {
+ StatusCode::OK => {}
+ StatusCode::NO_CONTENT => return Ok(false),
+ status => return Err(ClientErr::Sse(format!("expected HTTP 200, got {status}"))),
+ }
+ let content_type = self.str_header(header::CONTENT_TYPE)?;
+ if !content_type.eq_ignore_ascii_case("text/event-stream") {
+ return Err(ClientErr::Sse(format!(
+ "expected text/event-stream, got {content_type}"
+ )));
+ }
+ client.connect(self.body);
+ Ok(true)
}
async fn full_body(self) -> Result<Bytes, ClientErr> {
// Max 1 mb
- Limited::new(self.body, 1 * 1024 * 1024)
+ Limited::new(self.body, 1024 * 1024)
.collect()
.await
.map(|it| it.to_bytes())
@@ -308,3 +296,17 @@ impl Res {
Ok(parsed)
}
}
+
+impl HeaderParser<ClientErr> for Res {
+ fn parse<T>(
+ &self,
+ name: impl Into<HeaderName>,
+ kind: &'static str,
+ transform: impl FnOnce(&str) -> Result<T, ()>,
+ ) -> Result<T, ClientErr> {
+ self.head
+ .headers
+ .parse(name, kind, transform)
+ .map_err(ClientErr::Headers)
+ }
+}
diff --git a/common/http-client/src/headers.rs b/common/http-client/src/headers.rs
@@ -19,39 +19,40 @@ use hyper::{HeaderMap, header::HeaderName};
#[derive(Debug, thiserror::Error)]
pub enum HeaderError {
#[error("Missing header {0}")]
- Missing(&'static str),
+ Missing(HeaderName),
#[error("Malformed header {0} expected string got binary")]
- NotStr(&'static str),
+ NotStr(HeaderName),
#[error("Malformed header {0} expected {1} got '{2}'")]
- Malformed(&'static str, &'static str, Box<str>),
+ Malformed(HeaderName, &'static str, Box<str>),
}
-pub trait HeaderParser {
+pub trait HeaderParser<E> {
fn parse<T>(
&self,
- name: &'static str,
+ name: impl Into<HeaderName>,
kind: &'static str,
transform: impl FnOnce(&str) -> Result<T, ()>,
- ) -> Result<T, HeaderError>;
- fn str_header(&self, name: &'static str) -> Result<String, HeaderError> {
+ ) -> Result<T, E>;
+ fn str_header(&self, name: impl Into<HeaderName>) -> Result<String, E> {
self.parse(name, "string", |s| Ok(s.to_owned()))
}
- fn int_header(&self, name: &'static str) -> Result<u64, HeaderError> {
+ fn int_header(&self, name: impl Into<HeaderName>) -> Result<u64, E> {
self.parse(name, "integer", |s| s.parse().map_err(|_| ()))
}
- fn bool_header(&self, name: &'static str) -> Result<bool, HeaderError> {
+ fn bool_header(&self, name: impl Into<HeaderName>) -> Result<bool, E> {
self.parse(name, "boolean", |s| s.parse().map_err(|_| ()))
}
}
-impl HeaderParser for HeaderMap {
+impl HeaderParser<HeaderError> for HeaderMap {
fn parse<T>(
&self,
- name: &'static str,
+ name: impl Into<HeaderName>,
kind: &'static str,
transform: impl FnOnce(&str) -> Result<T, ()>,
) -> Result<T, HeaderError> {
- let Some(value) = self.get(HeaderName::from_static(name)) else {
+ let name = name.into();
+ let Some(value) = self.get(&name) else {
return Err(HeaderError::Missing(name));
};
let Ok(str) = value.to_str() else {
diff --git a/common/http-client/src/lib.rs b/common/http-client/src/lib.rs
@@ -56,6 +56,8 @@ pub enum ClientErr {
Headers(#[from] HeaderError),
#[error("response: {0}")]
ResTransport(String),
+ #[error("response SSE: {0}")]
+ Sse(String),
}
pub fn client() -> Client {
diff --git a/common/http-client/src/sse.rs b/common/http-client/src/sse.rs
@@ -14,102 +14,59 @@
TALER; see the file COPYING. If not, see <http://www.gnu.org/licenses/>
*/
-use std::pin::Pin;
+use std::{borrow::Cow, num::NonZeroUsize};
-use compact_str::CompactString;
-use futures_util::{Stream, StreamExt as _, stream};
-use tokio_util::{
- bytes::Bytes,
- codec::{FramedRead, LinesCodec, LinesCodecError},
- io::StreamReader,
-};
-use tracing::trace;
+use futures_util::StreamExt as _;
+use http_body_util::BodyDataStream;
+use hyper::body::Incoming;
+use sse_core::{SseDecoder, SseEvent, SseStream, SseStreamError};
#[derive(Debug, Default, PartialEq, Eq)]
pub struct SseMessage {
- pub event: CompactString,
+ pub event: Cow<'static, str>,
pub data: String,
}
-type SseStream = dyn Stream<Item = std::result::Result<String, LinesCodecError>> + Send;
-
/// Server-sent event client
pub struct SseClient {
- pub last_event_id: Option<CompactString>,
pub reconnection_time: Option<u64>,
- stream: Pin<Box<SseStream>>,
+ stream: sse_core::SseStream<BodyDataStream<Incoming>>,
}
impl SseClient {
pub fn new() -> Self {
Self {
- last_event_id: None,
reconnection_time: None,
- stream: Box::pin(stream::empty()),
+ stream: SseStream::with_decoder(SseDecoder::with_limit(
+ NonZeroUsize::new(1024 * 1024).unwrap(),
+ )),
}
}
- pub fn connect<E: std::error::Error + Send + Sync + 'static>(
- &mut self,
- stream: impl Stream<Item = Result<Bytes, E>> + 'static + Send,
- ) {
- let stream = stream.map(|it| it.map_err(std::io::Error::other));
- let lines = FramedRead::new(StreamReader::new(stream), LinesCodec::new());
- self.stream = Box::pin(lines);
+ pub fn last_event_id(&self) -> Option<&str> {
+ self.stream.last_event_id().map(|it| it.as_ref())
}
- pub async fn next(&mut self) -> Option<Result<SseMessage, LinesCodecError>> {
- // TODO add tests
- let mut event = CompactString::new("message");
- let mut data = None::<String>;
- while let Some(res) = self.stream.next().await {
- let line = match res {
- Ok(line) => line,
- Err(e) => return Some(Err(e)),
- };
- // Parse line
- let (field, value): (&str, &str) = if line.is_empty() {
- if let Some(data) = data.take() {
- return Some(Ok(SseMessage { event, data }));
- } else {
- event = CompactString::new("message");
- continue;
- }
- } else if let Some(comment) = line.strip_prefix(':') {
- trace!(target: "sse", "{comment}");
- continue;
- } else if let Some((k, v)) = line.split_once(':') {
- (k, v.strip_prefix(' ').unwrap_or(v))
- } else {
- (&line, "")
- };
+ pub fn connect(&mut self, body: Incoming) {
+ self.stream.attach(BodyDataStream::new(body));
+ }
- // Process field
- match field {
- "event" => event = CompactString::new(value),
- "data" => match data.as_mut() {
- Some(data) => {
- data.push('\n');
- data.push_str(value);
- }
- None => data = Some(value.to_string()),
- },
- "id" => {
- if !value.contains('\0') {
- self.last_event_id = Some(CompactString::new(value))
- }
+ pub async fn next(&mut self) -> Option<Result<SseMessage, SseStreamError<hyper::Error>>> {
+ loop {
+ match self.stream.next().await {
+ Some(Ok(SseEvent::Message(message))) => {
+ return Some(Ok(SseMessage {
+ event: message.event,
+ data: message.data,
+ }));
}
- "retry" => {
- if value.as_bytes().iter().all(|c| c.is_ascii_digit())
- && let Ok(int) = value.parse::<u64>()
- {
- self.reconnection_time = Some(int)
- }
+ Some(Ok(SseEvent::Retry(milliseconds))) => {
+ self.reconnection_time = Some(u64::from(milliseconds));
}
- _ => continue,
+ Some(Err(e)) => return Some(Err(e)),
+ None => return None,
}
}
- None
}
}
@@ -118,104 +75,3 @@ impl Default for SseClient {
Self::new()
}
}
-
-#[tokio::test]
-pub async fn protocol() {
- pub async fn test(
- stream: &'static str,
- result: &[(&str, &str)],
- last_event_id: Option<&str>,
- reconnection_time: Option<u64>,
- ) {
- let stream = stream::iter(
- stream
- .as_bytes()
- .chunks(12)
- .map(|chunk| std::io::Result::Ok(Bytes::from_static(chunk))),
- );
- let mut client = SseClient::new();
- client.connect(stream);
- let mut res = Vec::new();
- while let Some(msg) = client.next().await {
- res.push(msg.unwrap());
- }
- assert_eq!(
- result,
- &res.iter()
- .map(|m| (m.event.as_ref(), m.data.as_ref()))
- .collect::<Vec<_>>()
- );
- assert_eq!(client.last_event_id.as_deref(), last_event_id);
- assert_eq!(client.reconnection_time, reconnection_time);
- }
-
- macro_rules! check {
- // Handle multiple tuples + optional id and retry
- ($stream:expr $(, ($e:expr, $d:expr))* $(, id: $id:expr)? $(, retry: $retry:expr)?) => {
- test(
- $stream,
- &[ $( ($e, $d) ),* ],
- { let mut _id = None; $( _id = Some($id); )? _id },
- { let mut _r = None; $( _r = Some($retry); )? _r }
- ).await
- };
- }
-
- check!("data\n\n", ("message", ""));
- check!("data:key:value\n\n", ("message", "key:value"));
- check!("data: value\n\n", ("message", "value"));
- check!("data:first\ndata:second\n\n", ("message", "first\nsecond"));
-
- check!(
- "event:first\nevent:second\ndata:test\n\n",
- ("second", "test")
- );
-
- check!("data:test\r\n\r\n", ("message", "test"));
- check!("data:test\r\r");
- check!("data:line1\r\ndata:line2\n\n", ("message", "line1\nline2"));
- check!("data:test\n");
- check!("data:test\n\n\n", ("message", "test"));
- check!("data:\ndata:\n\n", ("message", "\n"));
- check!("data:\ndata:content\ndata:\n\n", ("message", "\ncontent\n"));
- check!("data:Hello δΈη π\n\n", ("message", "Hello δΈη π"));
- check!("data: \n\n", ("message", " "));
- check!("id:123\ndata:test\n\n", ("message", "test"), id: "123");
- check!("id:first\nid:second\ndata:test\n\n", ("message", "test"), id: "second");
- check!("id:first\nid:second\nid:\ndata:test\n\n", ("message", "test"), id: "");
- check!("id:\ndata:test\n\n", ("message", "test"), id: "");
- check!(
- "id:test:123\ndata:test\n\n",
- ("message", "test"),
- id: "test:123"
- );
- check!("id:123\x00456\ndata:test\n\n", ("message", "test"));
- check!("id:123\n\n", id: "123");
- check!(
- "id:123\ndata:first\n\ndata:second\n\n",
- ("message", "first"), ("message", "second"),
- id: "123"
- );
- check!("event:customEvent\ndata:test\n\n", ("customEvent", "test"));
- check!("event:\ndata:test\n\n", ("", "test"));
- check!("event:my event\ndata:test\n\n", ("my event", "test"));
-
- check!("retry:3000\n\n", retry: 3000);
- check!("retry:0\n\n", retry: 0);
- check!("retry:abc\n\n");
- check!("retry:-1000\n\n");
- check!("retry:1000.5\n\n");
- check!("retry:1000\nretry:2000\n\n", retry: 2000);
-
- check!(":comment\n\n");
- check!(": comment\n\n");
- check!(":comment\ndata:test\n\n", ("message", "test"));
- check!("unknown:value\ndata:test\n\n", ("message", "test"));
- check!("datta:test\n\n");
- check!(" data:test\n\n");
- check!("data :test\n\n");
- check!("data:\tvalue\n\n", ("message", "\tvalue"));
- check!("id:123\n\n", id: "123");
- check!("event:test\n\n");
- check!("event:test\n\ndata:value\n\n", ("message", "value"));
-}
diff --git a/common/taler-api/src/db.rs b/common/taler-api/src/db.rs
@@ -151,14 +151,13 @@ pub async fn pooling<R, N, F: Future<Output = sqlx::Result<R>>>(
let init = load().await?;
// Long polling if we found no transactions
if !check(&init) {
- let pooling = tokio::time::timeout(Duration::from_millis(timeout), async {
+ tokio::time::timeout(Duration::from_millis(timeout), async {
listener.wait_for(filter).await.ok();
})
- .await;
- match pooling {
- Ok(_) => load().await,
- Err(_) => Ok(init),
- }
+ .await
+ .ok();
+ // Whether pooling worked or not we load some fresh data from the database
+ load().await
} else {
Ok(init)
}