use axum::{
body::Body,
extract::{RawQuery, State},
http::{header, HeaderMap, Response, StatusCode},
response::{
sse::{Event, KeepAlive, Sse},
IntoResponse,
},
};
use futures_util::stream::Stream;
use pubky_common::crypto::PublicKey;
use serde::Deserialize;
use std::{collections::HashMap, convert::Infallible, time::Instant};
use tower_cookies::Cookies;
use url::form_urlencoded;
use crate::{
client_server::{
auth::{
grant::bearer::extract_bearer_token, has_read_permission, AuthSession,
PendingStreamAuth,
},
query_params::ListQueryParams,
AppState,
},
constants::{PRIVATE_ROOT, PUBLIC_ROOT},
observability::ConnectionGuard,
persistence::{
files::events::{
EventCursor, EventEntity, EventsService, PathFilter, MAX_EVENT_STREAM_USERS,
},
sql::SqlDb,
},
shared::{webdav::StoragePath, HttpError, HttpResult},
};
#[derive(Debug, thiserror::Error)]
pub enum EventStreamError {
#[error("User not found")]
UserNotFound,
#[error("{0}")]
InvalidParameter(String),
#[error("Database error: {0}")]
DatabaseError(#[from] sqlx::Error),
#[error("Invalid public key: {0}")]
InvalidPublicKey(String),
}
impl From<EventStreamError> for HttpError {
fn from(error: EventStreamError) -> Self {
match error {
EventStreamError::UserNotFound => HttpError::not_found(),
EventStreamError::DatabaseError(e) => HttpError::from(e),
_ => HttpError::bad_request(error.to_string()),
}
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(try_from = "RawEventStreamQueryParams")]
pub struct EventStreamQueryParams {
pub limit: Option<u16>,
pub reverse: bool,
pub live: bool,
pub user_cursors: Vec<(PublicKey, Option<String>)>,
pub paths: Vec<StoragePath>,
}
#[derive(Clone, Copy)]
enum EventStreamTenantScope<'a> {
PublicOnly,
PrivateSingleTenant(&'a PublicKey),
PrivateUnsupported,
}
impl<'a> EventStreamTenantScope<'a> {
fn from_query(paths: &[StoragePath], user_cursors: &'a [(PublicKey, Option<String>)]) -> Self {
if !paths
.iter()
.any(|path| path.as_str().starts_with(PRIVATE_ROOT))
{
return Self::PublicOnly;
}
match user_cursors {
[(tenant, _)] => Self::PrivateSingleTenant(tenant),
_ => Self::PrivateUnsupported,
}
}
fn tenant(self) -> Option<&'a PublicKey> {
match self {
Self::PrivateSingleTenant(tenant) => Some(tenant),
Self::PublicOnly | Self::PrivateUnsupported => None,
}
}
}
#[derive(Debug, Deserialize)]
struct RawEventStreamQueryParams {
#[serde(default)]
user: Vec<String>,
limit: Option<u16>,
#[serde(default)]
reverse: bool,
#[serde(default)]
live: bool,
#[serde(default)]
paths: Vec<String>,
}
fn parse_query_params(query: &str) -> Result<EventStreamQueryParams, EventStreamError> {
let mut users = Vec::new();
let mut limit = None;
let mut reverse = false;
let mut live = false;
let mut paths = Vec::new();
for (key, value) in form_urlencoded::parse(query.as_bytes()) {
match key.as_ref() {
"user" => users.push(value.to_string()),
"limit" => {
let parsed = value.parse::<u16>().map_err(|_| {
EventStreamError::InvalidParameter(format!("Invalid limit: {}", value))
})?;
if parsed == 0 {
return Err(EventStreamError::InvalidParameter(
"limit must be at least 1".to_string(),
));
}
limit = Some(parsed);
}
"reverse" => {
reverse = value == "true" || value == "1";
}
"live" => {
live = value == "true" || value == "1";
}
"path" if !value.is_empty() => {
paths.push(value.to_string());
}
_ => {} }
}
let raw = RawEventStreamQueryParams {
user: users,
limit,
reverse,
live,
paths,
};
raw.try_into()
}
impl TryFrom<RawEventStreamQueryParams> for EventStreamQueryParams {
type Error = EventStreamError;
fn try_from(raw: RawEventStreamQueryParams) -> Result<Self, Self::Error> {
if raw.live && raw.reverse {
return Err(EventStreamError::InvalidParameter(
"Cannot use live mode with reverse ordering".to_string(),
));
}
let mut user_cursors = Vec::new();
for value in raw.user {
if value.is_empty() {
continue;
}
let (pubkey_str, cursor_str) = if let Some((pubkey, cursor)) = value.split_once(':') {
(pubkey, Some(cursor))
} else {
(value.as_str(), None)
};
if PublicKey::is_pubky_prefixed(pubkey_str) {
return Err(EventStreamError::InvalidPublicKey(pubkey_str.to_string()));
}
let pubkey = PublicKey::try_from_z32(pubkey_str)
.map_err(|_| EventStreamError::InvalidPublicKey(pubkey_str.to_string()))?;
user_cursors.push((pubkey, cursor_str.map(|s| s.to_string())));
}
if user_cursors.is_empty() {
return Err(EventStreamError::InvalidParameter(
"user parameter is required".to_string(),
));
}
if user_cursors.len() > MAX_EVENT_STREAM_USERS {
return Err(EventStreamError::InvalidParameter(format!(
"Too many users. Maximum allowed: {}",
MAX_EVENT_STREAM_USERS
)));
}
let mut paths = Vec::with_capacity(raw.paths.len());
for p in raw.paths {
let normalized_path = if p.starts_with('/') {
p
} else {
format!("/{}", p)
};
let path = StoragePath::normalize(&normalized_path).map_err(|_| {
EventStreamError::InvalidParameter(format!("Invalid path: {}", normalized_path))
})?;
paths.push(path);
}
Ok(EventStreamQueryParams {
limit: raw.limit,
reverse: raw.reverse,
live: raw.live,
user_cursors,
paths,
})
}
}
fn format_events_feed(events: &[EventEntity]) -> String {
let mut result = events
.iter()
.map(|event| format!("{} {}", event.event_type, event.pubky_uri()))
.collect::<Vec<String>>();
if let Some(next_cursor) = events.last().map(|event| event.id.to_string()) {
result.push(format!("cursor: {}", next_cursor));
}
result.join("\n")
}
pub async fn feed(
State(state): State<AppState>,
params: ListQueryParams,
) -> HttpResult<impl IntoResponse> {
let cursor = match params.cursor {
Some(cursor) => cursor,
None => "0".to_string(),
};
let cursor = match state
.events_service
.parse_cursor(cursor.as_str(), &mut state.sql_db.pool().into())
.await
{
Ok(cursor) => cursor,
Err(_e) => return Err(HttpError::bad_request("Invalid cursor")),
};
let query_start = Instant::now();
let events = state
.events_service
.get_public_by_cursor(Some(cursor), params.limit, &mut state.sql_db.pool().into())
.await?;
state
.metrics
.record_events_db_query(query_start.elapsed().as_millis());
Ok(Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, "text/plain")
.body(Body::from(format_events_feed(&events)))
.unwrap())
}
pub async fn feed_stream(
State(state): State<AppState>,
session: Option<AuthSession>,
headers: HeaderMap,
cookies: Cookies,
raw_query: RawQuery,
) -> HttpResult<Sse<impl Stream<Item = Result<Event, Infallible>>>> {
let params =
parse_query_params(raw_query.0.as_deref().unwrap_or("")).map_err(HttpError::from)?;
let tenant_scope = EventStreamTenantScope::from_query(¶ms.paths, ¶ms.user_cursors);
let session = match session {
Some(AuthSession::Cookie(_)) if has_bearer_auth(&headers) => None,
Some(session) => Some(session),
None if !has_bearer_auth(&headers) => {
resolve_tenant_cookie_session(&state, &cookies, tenant_scope).await
}
None => None,
};
let allowed_paths = authorized_paths(¶ms.paths, tenant_scope, session.as_ref())?;
let pending_stream_auth = PendingStreamAuth::subscribe(
matches!(tenant_scope, EventStreamTenantScope::PrivateSingleTenant(_)),
&state.auth_state,
)
.await
.map_err(|_| {
HttpError::new_with_message(
StatusCode::SERVICE_UNAVAILABLE,
"Private event streams are temporarily unavailable",
)
})?;
let mut stream_auth = pending_stream_auth
.authorize(session, &state.auth_state)
.await?;
let mut user_cursor_map =
resolve_user_cursors(¶ms.user_cursors, &state.events_service, &state.sql_db)
.await
.map_err(HttpError::from)?;
let mut total_sent: usize = 0;
let stream = async_stream::stream! {
let _guard = ConnectionGuard::new(state.metrics.clone());
let mut rx = state.events_service.subscribe();
loop {
if !stream_auth.is_valid(&state.auth_state).await {
return;
}
while rx.try_recv().is_ok() {}
let current_user_cursors: Vec<(i32, Option<EventCursor>)> =
user_cursor_map.iter().map(|(k, cursor)| (*k, *cursor)).collect();
let query_start = Instant::now();
let events = match state
.events_service
.get_by_user_cursors(
current_user_cursors,
params.reverse,
&allowed_paths,
&mut state.sql_db.pool().into(),
)
.await
{
Ok(events) => events,
Err(e) => {
tracing::error!("Database error while fetching events: {}", e);
break;
}
};
state.metrics.record_event_stream_db_query(query_start.elapsed().as_millis());
let event_count = events.len();
for event in events {
if !stream_auth.is_valid(&state.auth_state).await {
return;
}
user_cursor_map.insert(event.user_id, Some(event.cursor()));
yield Ok(Event::default()
.event(event.event_type.to_string())
.data(event.to_sse_data()));
total_sent += 1;
if let Some(max) = params.limit {
if total_sent >= max as usize {
return;
}
}
}
if event_count == 0 {
if !params.live {
return;
}
break;
}
}
if params.live {
let user_ids: Vec<i32> = user_cursor_map.keys().copied().collect();
let half_capacity = state.events_service.channel_capacity() / 2;
loop {
tokio::select! {
biased;
check = stream_auth.next_check(&state.auth_state) => {
if !check.await {
return;
}
}
event = rx.recv() => match event {
Ok(event) => {
if rx.len() >= half_capacity {
state.metrics.record_broadcast_half_full();
}
if !should_include_live_event(&event, &user_ids, &user_cursor_map, &allowed_paths) {
continue;
}
user_cursor_map.insert(event.user_id, Some(event.cursor()));
yield Ok(Event::default()
.event(event.event_type.to_string())
.data(event.to_sse_data()));
total_sent += 1;
if let Some(max) = params.limit {
if total_sent >= max as usize {
return;
}
}
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => {
state.metrics.record_broadcast_lagged();
tracing::warn!(
"Slow client detected: broadcast channel lagged by {} events. Closing connection.",
skipped
);
return;
}
Err(_) => break, }
}
}
}
};
Ok(Sse::new(stream).keep_alive(KeepAlive::default()))
}
async fn resolve_user_cursors(
user_cursors: &[(PublicKey, Option<String>)],
events_service: &EventsService,
sql_db: &SqlDb,
) -> Result<HashMap<i32, Option<EventCursor>>, EventStreamError> {
use crate::persistence::sql::user::UserRepository;
let mut user_cursor_map: HashMap<i32, Option<EventCursor>> = HashMap::new();
for (user_pubkey, cursor_str_opt) in user_cursors {
let user_id = UserRepository::get_id(user_pubkey, &mut sql_db.pool().into())
.await
.map_err(|e| match e {
sqlx::Error::RowNotFound => EventStreamError::UserNotFound,
e => EventStreamError::DatabaseError(e),
})?;
let cursor = if let Some(cursor_str) = cursor_str_opt {
Some(
events_service
.parse_cursor(cursor_str, &mut sql_db.pool().into())
.await
.map_err(|_| {
EventStreamError::InvalidParameter(format!(
"Invalid cursor: {}",
cursor_str
))
})?,
)
} else {
None
};
user_cursor_map.insert(user_id, cursor);
}
Ok(user_cursor_map)
}
fn has_bearer_auth(headers: &HeaderMap) -> bool {
extract_bearer_token(headers).has_bearer_scheme()
}
async fn resolve_tenant_cookie_session(
state: &AppState,
cookies: &Cookies,
scope: EventStreamTenantScope<'_>,
) -> Option<AuthSession> {
let EventStreamTenantScope::PrivateSingleTenant(tenant) = scope else {
return None;
};
let cookie_value = cookies.get(&tenant.z32()).map(|c| c.value().to_string());
state
.auth_state
.cookie_auth_service
.resolve_session_from_cookie(cookie_value, tenant)
.await
}
fn authorized_paths(
paths: &[StoragePath],
scope: EventStreamTenantScope<'_>,
session: Option<&AuthSession>,
) -> Result<Vec<PathFilter>, HttpError> {
if paths.is_empty() {
return Ok(vec![StoragePath::new(PUBLIC_ROOT)
.expect("public root is canonical")
.into()]);
}
let mut allowed = Vec::with_capacity(paths.len());
for path in paths {
has_read_permission(session, scope.tenant(), path)?;
allowed.push(path.clone().into());
}
Ok(allowed)
}
fn should_include_live_event(
event: &EventEntity,
user_ids: &[i32],
user_cursor_map: &HashMap<i32, Option<EventCursor>>,
allowed_paths: &[PathFilter],
) -> bool {
if !user_ids.contains(&event.user_id) {
return false;
}
if let Some(Some(cursor)) = user_cursor_map.get(&event.user_id) {
if event.cursor() <= *cursor {
return false;
}
}
let path = event.path.path().as_str();
allowed_paths.iter().any(|filter| filter.matches(path))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::client_server::auth::cookie::persistence::{SessionEntity, SessionSecret};
use crate::client_server::auth::grant::session::GrantSession;
use pubky_common::auth::jws::GrantId;
use pubky_common::capabilities::{Capabilities, Capability};
use pubky_common::crypto::Keypair;
fn pk() -> PublicKey {
Keypair::random().public_key()
}
fn wd(s: &str) -> StoragePath {
StoragePath::new(s).expect("valid test path")
}
fn pf(s: &str) -> PathFilter {
wd(s).into()
}
fn grant_session(user_key: PublicKey, capabilities: Capabilities) -> AuthSession {
AuthSession::Grant(GrantSession::test(
user_key,
capabilities,
GrantId::generate(),
9999999999,
))
}
fn cookie_session(user_key: PublicKey, capabilities: Capabilities) -> AuthSession {
AuthSession::Cookie(SessionEntity {
id: 1,
secret: SessionSecret::random(),
user_id: 1,
user_pubkey: user_key,
capabilities,
created_at: sqlx::types::chrono::DateTime::from_timestamp(0, 0)
.expect("valid timestamp")
.naive_utc(),
})
}
fn cursors(keys: &[&PublicKey]) -> Vec<(PublicKey, Option<String>)> {
keys.iter().map(|k| ((*k).clone(), None)).collect()
}
fn reject_status(result: Result<Vec<PathFilter>, HttpError>) -> StatusCode {
result
.expect_err("expected the subscription to be rejected")
.into_response()
.status()
}
fn authorize(
paths: &[StoragePath],
user_cursors: &[(PublicKey, Option<String>)],
session: Option<&AuthSession>,
) -> Result<Vec<PathFilter>, HttpError> {
authorized_paths(
paths,
EventStreamTenantScope::from_query(paths, user_cursors),
session,
)
}
fn authorization_headers(value: axum::http::HeaderValue) -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(header::AUTHORIZATION, value);
headers
}
#[test]
fn has_bearer_auth_uses_bearer_scheme() {
let cases = [
(b"Bearer token".as_slice(), true),
(b"Bearer \xff".as_slice(), true),
(b"Basic \xff".as_slice(), false),
];
for (value, expected) in cases {
let headers = authorization_headers(
axum::http::HeaderValue::from_bytes(value).expect("valid header bytes"),
);
assert_eq!(has_bearer_auth(&headers), expected, "{value:?}");
}
}
#[test]
fn parse_repeated_paths_preserves_order_and_trailing_slash() {
let q = format!(
"user={}&path=/pub/&path=/priv/app/&path=/priv/file",
pk().z32()
);
let params = parse_query_params(&q).unwrap();
let strs: Vec<&str> = params.paths.iter().map(|p| p.as_str()).collect();
assert_eq!(strs, vec!["/pub/", "/priv/app/", "/priv/file"]);
}
#[test]
fn parse_ignores_empty_path_values() {
let params =
parse_query_params(&format!("user={}&path=&path=/pub/&path=", pk().z32())).unwrap();
assert_eq!(
params.paths.iter().map(|p| p.as_str()).collect::<Vec<_>>(),
vec!["/pub/"]
);
}
#[test]
fn parse_requires_user() {
let err = parse_query_params("path=/pub/").unwrap_err();
assert_eq!(
HttpError::from(err).into_response().status(),
StatusCode::BAD_REQUEST
);
}
#[test]
fn parse_rejects_zero_limit() {
let err = parse_query_params(&format!("user={}&limit=0", pk().z32())).unwrap_err();
assert_eq!(err.to_string(), "limit must be at least 1");
}
#[test]
fn authorized_paths_defaults_to_public_dir_filter() {
let u = pk();
let filters = authorize(&[], &cursors(&[&u]), None).unwrap();
assert_eq!(filters, vec![pf("/pub/")]);
}
#[test]
fn authorized_paths_rejects_anonymous_private_path() {
let u = pk();
let status = reject_status(authorize(&[wd("/priv/app/")], &cursors(&[&u]), None));
assert_eq!(status, StatusCode::UNAUTHORIZED);
}
#[test]
fn authorized_paths_allows_cookie_session_own_private_path() {
let owner = pk();
let session = cookie_session(owner.clone(), Capabilities::from(vec![Capability::root()]));
let filters = authorize(&[wd("/priv/app/")], &cursors(&[&owner]), Some(&session)).unwrap();
assert_eq!(filters, vec![pf("/priv/app/")]);
}
#[test]
fn authorized_paths_rejects_wrong_tenant() {
let (a, b) = (pk(), pk());
let session = cookie_session(a, Capabilities::from(vec![Capability::root()]));
let status = reject_status(authorize(
&[wd("/priv/app/")],
&cursors(&[&b]),
Some(&session),
));
assert_eq!(status, StatusCode::FORBIDDEN);
}
#[test]
fn authorized_paths_rejects_private_path_with_multiple_users() {
let (a, b) = (pk(), pk());
let session = grant_session(a.clone(), Capabilities::from(vec![Capability::root()]));
let status = reject_status(authorize(
&[wd("/priv/app/")],
&cursors(&[&a, &b]),
Some(&session),
));
assert_eq!(status, StatusCode::FORBIDDEN);
}
#[test]
fn authorized_paths_rejects_under_scoped_private_path() {
let owner = pk();
let session = grant_session(
owner.clone(),
Capabilities::from(vec![Capability::read("/priv/app/").unwrap()]),
);
let status = reject_status(authorize(
&[wd("/priv/other/")],
&cursors(&[&owner]),
Some(&session),
));
assert_eq!(status, StatusCode::FORBIDDEN);
}
#[test]
fn authorized_paths_allows_mixed_public_and_private_union() {
let owner = pk();
let session = grant_session(
owner.clone(),
Capabilities::from(vec![Capability::read("/priv/app/").unwrap()]),
);
let filters = authorize(
&[wd("/pub/"), wd("/priv/app/")],
&cursors(&[&owner]),
Some(&session),
)
.unwrap();
assert_eq!(filters, vec![pf("/pub/"), pf("/priv/app/")]);
}
#[test]
fn authorized_paths_allows_public_paths_with_multiple_users() {
let (a, b) = (pk(), pk());
let filters = authorize(&[wd("/pub/")], &cursors(&[&a, &b]), None).unwrap();
assert_eq!(filters, vec![pf("/pub/")]);
}
}