taler-rust

GNU Taler code in Rust. Largely core banking integrations.
Log | Files | Refs | Submodules | README | LICENSE

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:
Madapters/taler-cyclos/src/cyclos_api/api.rs | 7++++---
Madapters/taler-cyclos/src/cyclos_api/client.rs | 2+-
Madapters/taler-cyclos/src/notification.rs | 17++++++++++++-----
Madapters/taler-cyclos/src/worker.rs | 1-
Madapters/taler-magnet-bank/src/magnet_api/api.rs | 3++-
Madapters/taler-wise/src/wise_api/client.rs | 4++--
Mcommon/http-client/Cargo.toml | 1+
Mcommon/http-client/src/builder.rs | 60+++++++++++++++++++++++++++++++-----------------------------
Mcommon/http-client/src/headers.rs | 25+++++++++++++------------
Mcommon/http-client/src/lib.rs | 2++
Mcommon/http-client/src/sse.rs | 198+++++++++++--------------------------------------------------------------------
Mcommon/taler-api/src/db.rs | 11+++++------
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) }