Skip to content

Commit d7a91a2

Browse files
committed
shared conn for listen
1 parent c0b88f6 commit d7a91a2

5 files changed

Lines changed: 77 additions & 34 deletions

File tree

Cargo.lock

Lines changed: 1 addition & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

Cargo.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -54,7 +54,7 @@ tokio-postgres-rustls = "^0"
5454
rustls = "^0"
5555
axum-server = { version = "^0", features = ["tls-rustls"] }
5656
futures = "^0"
57-
tokio-stream = "^0"
57+
tokio-stream = { version = "^0", features = ["sync"] }
5858
serde_qs = "1.0.0-rc.4"
5959
bytes = "^1"
6060
sqlparser = { version = "^0", features = ["visitor"] }

src/error/mod.rs

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,3 @@
1-
use std::{backtrace::Backtrace, borrow::Cow};
2-
31
use axum::{extract::multipart, http::{self, header}, response::{IntoResponse, Response}};
42
use biscuit_auth::error;
53
use deadpool_postgres::{CreatePoolError, PoolError};

src/extract/query.rs

Lines changed: 17 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -371,11 +371,14 @@ fn parse_sql(order: &Option<BTreeMap<String, serde_json::Value>>, root_sql: Opti
371371
#[allow(clippy::unwrap_used)]
372372
#[allow(clippy::indexing_slicing)]
373373
mod tests {
374-
use axum::{body::Body, http::{header::CONTENT_TYPE, Request}};
374+
use std::sync::Arc;
375+
376+
use axum::{body::Body, http::{header::CONTENT_TYPE, Request}};
375377

376378
use axum::extract::FromRequest;
377379
use conf::Conf;
378-
use crate::extract::query::{Param, Query};
380+
use tokio_postgres::AsyncMessage;
381+
use crate::{extract::query::{Param, Query}};
379382

380383
#[tokio::test]
381384
async fn test_json_body() {
@@ -388,11 +391,17 @@ mod tests {
388391

389392
let read_pool = httpg_config.pg.read_pool().unwrap();
390393
let write_pool = httpg_config.pg.write_pool().unwrap();
394+
395+
let (client, mut _conn) = httpg_config.pg.connect().await.unwrap();
396+
397+
let (tx, _rx) = tokio::sync::broadcast::channel::<AsyncMessage>(16);
391398

392399
let state = crate::AppState {
393400
read_pool,
394401
write_pool,
395402
config: httpg_config.to_owned(),
403+
tx,
404+
client: Arc::new(client),
396405
};
397406
let q = Query::from_request(req, &state).await.unwrap();
398407

@@ -413,10 +422,16 @@ mod tests {
413422
let read_pool = httpg_config.pg.read_pool().unwrap();
414423
let write_pool = httpg_config.pg.write_pool().unwrap();
415424

425+
let (client, mut _conn) = httpg_config.pg.connect().await.unwrap();
426+
427+
let (tx, _rx) = tokio::sync::broadcast::channel::<AsyncMessage>(16);
428+
416429
let state = crate::AppState {
417430
read_pool,
418431
write_pool,
419432
config: httpg_config.to_owned(),
433+
tx,
434+
client: Arc::new(client),
420435
};
421436
let q = Query::from_request(req, &state).await.unwrap();
422437

src/main.rs

Lines changed: 58 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -18,15 +18,17 @@ use conf::Conf;
1818
use cookie::time::{Duration, OffsetDateTime};
1919
use futures::{StreamExt, TryStreamExt};
2020
use lettre::{
21-
AsyncSmtpTransport, AsyncTransport, Message, Tokio1Executor, message::header::ContentType, transport::smtp::authentication::{Credentials, Mechanism::Login}
21+
AsyncSmtpTransport, AsyncTransport, Message, Tokio1Executor, message::header::ContentType, transport::smtp::authentication::{Credentials}
2222
};
2323
use serde_json::json;
24+
use tokio::sync::{broadcast::Sender};
25+
use tokio_stream::wrappers::errors::BroadcastStreamRecvError;
2426
use tower::builder::ServiceBuilder;
25-
use tower_http::{compression::CompressionLayer, cors::{Any, CorsLayer}, services::ServeDir, trace::TraceLayer};
27+
use tower_http::{cors::{Any, CorsLayer}, services::ServeDir, trace::TraceLayer};
2628
use web_push::{ContentEncoding, HyperWebPushClient, SubscriptionInfo, VapidSignatureBuilder, WebPushClient, WebPushMessageBuilder};
27-
use std::{env, fs::{self, File}, net::{SocketAddr, TcpListener}};
29+
use std::{env, fs::{self, File}, net::{SocketAddr, TcpListener}, ops::Deref, sync::Arc};
2830
use std::collections::HashMap;
29-
use tokio_postgres::{IsolationLevel, types::{ToSql, Type}};
31+
use tokio_postgres::{AsyncMessage, Client, IsolationLevel, types::{ToSql, Type}};
3032
use deadpool_postgres::{Pool, Transaction};
3133
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
3234
use biscuit_auth::{KeyPair, PrivateKey, Biscuit, builder::*};
@@ -77,6 +79,8 @@ struct AppState {
7779
read_pool: Pool,
7880
write_pool: Pool,
7981
config: HttpgConfig,
82+
tx: Sender<AsyncMessage>,
83+
client: Arc<Client>,
8084
}
8185

8286
#[tokio::main]
@@ -93,11 +97,39 @@ async fn main() -> Result<(), HttpgError> {
9397

9498
let read_pool = httpg_config.pg.read_pool()?;
9599
let write_pool = httpg_config.pg.write_pool()?;
100+
101+
let (client, mut conn) = httpg_config.pg.connect().await?;
102+
103+
let (tx, _rx) = tokio::sync::broadcast::channel::<AsyncMessage>(16);
104+
105+
let mut stream = futures::stream::poll_fn(move |cx| conn.poll_message(cx));
106+
107+
let wrapped_tx = tx.clone();
108+
tokio::spawn(async move {
109+
while let Some(Ok(m)) = stream.next().await {
110+
wrapped_tx.send(m).unwrap();
111+
// Event::default().data(n.payload())
112+
// )).map_err(|e| HttpgError::anyhow(e.to_string()))?;
113+
// match m {
114+
// tokio_postgres::AsyncMessage::Notice(n) => tracing::info!("{n:#?}"),
115+
// tokio_postgres::AsyncMessage::Notification(n) => {
116+
// tx.send(Ok(
117+
// Event::default().data(n.payload())
118+
// )).map_err(|e| HttpgError::anyhow(e.to_string()))?;
119+
// },
120+
// _ => {HttpgError::anyhow("unsupported AsyncMessage");},
121+
// }
122+
}
123+
124+
Ok::<_, HttpgError>(())
125+
});
96126

97127
let state = AppState {
98128
read_pool,
99129
write_pool,
100130
config: httpg_config.to_owned(),
131+
tx,
132+
client: Arc::new(client),
101133
};
102134

103135
let app = Router::new()
@@ -309,8 +341,7 @@ async fn pre<'a>(tx: &mut Transaction<'a>, biscuit: &Option<extract::biscuit::Bi
309341
serde_json::to_string(&query)?,
310342
Type::TEXT
311343
)
312-
])
313-
.await?;
344+
]).await?;
314345

315346
if let Some(extract::biscuit::Biscuit(b)) = biscuit {
316347
futures::future::join_all(b.iter().map(async |sql| {
@@ -534,34 +565,32 @@ async fn stream_query(
534565

535566
#[debug_handler]
536567
async fn sse_query(
537-
State(AppState {config, ..}): State<AppState>,
538-
biscuit: Option<extract::biscuit::Biscuit>,
568+
State(AppState {tx, client, ..}): State<AppState>,
539569
Path(channel): Path<String>,
540-
query: extract::query::Query,
541570
) -> Result<impl IntoResponse, HttpgError> {
542571

543-
let (client, mut conn) = config.pg.connect().await?;
544-
545-
let (tx, rx) = tokio::sync::mpsc::unbounded_channel::<Result<Event, HttpgError>>();
546-
547-
let mut stream = futures::stream::poll_fn(move |cx| conn.poll_message(cx));
548-
tokio::spawn(async move {
549-
client.simple_query_raw(&format!("listen {channel}")).await.unwrap();
550-
while let Some(Ok(m)) = stream.next().await {
551-
match m {
552-
tokio_postgres::AsyncMessage::Notice(n) => tracing::info!("{n:#?}"),
553-
tokio_postgres::AsyncMessage::Notification(n) => {
554-
tx.send(Ok(
555-
Event::default().data(n.payload())
556-
)).unwrap();
557-
},
558-
_ => todo!()
559-
}
560-
}
561-
});
572+
client.simple_query_raw(&format!("listen {channel}")).await?;
562573

563574
Ok(Sse::new(
564-
tokio_stream::wrappers::UnboundedReceiverStream::new(rx)
575+
tokio_stream::wrappers::BroadcastStream::new(tx.subscribe())
576+
.try_filter_map(move |m| {
577+
let channel = channel.clone();
578+
async move {
579+
match m {
580+
tokio_postgres::AsyncMessage::Notice(n) => {
581+
tracing::info!("{n:#?}");
582+
Ok(None)
583+
},
584+
tokio_postgres::AsyncMessage::Notification(n) => {
585+
match n.channel() {
586+
c if c == channel => Ok(Some(Event::default().data(n.payload()))),
587+
_ => Ok(None)
588+
}
589+
},
590+
_ => Err(BroadcastStreamRecvError::Lagged(1)),
591+
}
592+
}
593+
})
565594
))
566595
// .keep_alive(
567596
// axum::response::sse::KeepAlive::new()

0 commit comments

Comments
 (0)