use std::convert::Infallible;
use axum::{
extract::{RawQuery, State},
http::{header, HeaderValue},
response::{
sse::{Event, KeepAlive, Sse},
IntoResponse,
},
};
use futures_util::StreamExt;
use pubky_common::crypto::PublicKey;
use url::form_urlencoded;
use super::super::app_state::AppState;
use crate::{
persistence::{
files::events::{AllEventsFilter, Mode, PathFilter, MAX_EVENT_STREAM_USERS},
sql::user::UserRepository,
},
shared::{webdav::StoragePath, HttpError, HttpResult},
};
struct AdminStreamParams {
users: Vec<PublicKey>,
cursor: Option<String>,
limit: Option<u16>,
mode: Mode,
paths: Vec<StoragePath>,
}
fn parse_admin_stream_query(query: &str) -> Result<AdminStreamParams, HttpError> {
let mut users: Vec<PublicKey> = Vec::new();
let mut cursor = None;
let mut limit = None;
let mut live = false;
let mut reverse = false;
let mut paths: Vec<StoragePath> = Vec::new();
for (key, value) in form_urlencoded::parse(query.as_bytes()) {
match key.as_ref() {
"user" => {
if value.is_empty() {
continue;
}
if PublicKey::is_pubky_prefixed(value.as_ref()) {
return Err(HttpError::bad_request(format!(
"Invalid public key: {value}"
)));
}
let pk = PublicKey::try_from_z32(value.as_ref())
.map_err(|_| HttpError::bad_request(format!("Invalid public key: {value}")))?;
if !users.contains(&pk) {
users.push(pk);
}
}
"cursor" => {
if !value.is_empty() {
cursor = Some(value.to_string());
}
}
"limit" => {
let parsed = value
.parse::<u16>()
.map_err(|_| HttpError::bad_request(format!("Invalid limit: {value}")))?;
if parsed == 0 {
return Err(HttpError::bad_request("limit must be at least 1"));
}
limit = Some(parsed);
}
"live" => live = value == "true" || value == "1",
"reverse" => reverse = value == "true" || value == "1",
"path" => {
if value.is_empty() {
continue;
}
let normalized = if value.starts_with('/') {
value.into_owned()
} else {
format!("/{value}")
};
let path = StoragePath::normalize(&normalized)
.map_err(|_| HttpError::bad_request(format!("Invalid path: {normalized}")))?;
paths.push(path);
}
_ => {} }
}
let mode = match (live, reverse) {
(false, false) => Mode::Forward,
(true, false) => Mode::ForwardLive,
(false, true) => Mode::Reverse,
(true, true) => {
return Err(HttpError::bad_request(
"Cannot use live mode with reverse ordering",
))
}
};
if users.len() > MAX_EVENT_STREAM_USERS {
return Err(HttpError::bad_request(format!(
"Too many users. Maximum allowed: {MAX_EVENT_STREAM_USERS}"
)));
}
Ok(AdminStreamParams {
users,
cursor,
limit,
mode,
paths,
})
}
async fn resolve_filter(
state: &AppState,
params: AdminStreamParams,
) -> HttpResult<AllEventsFilter> {
let user_ids = if params.users.is_empty() {
None
} else {
let mut ids = Vec::with_capacity(params.users.len());
for pk in ¶ms.users {
let id = UserRepository::get_id(pk, &mut state.sql_db.pool().into())
.await
.map_err(|e| match e {
sqlx::Error::RowNotFound => HttpError::not_found(),
e => HttpError::from(e),
})?;
ids.push(id);
}
Some(ids)
};
let start_cursor = match params.cursor.as_deref() {
Some(c) => Some(
state
.events_service
.parse_cursor(c, &mut state.sql_db.pool().into())
.await
.map_err(|_| HttpError::bad_request("Invalid cursor"))?,
),
None => None,
};
Ok(AllEventsFilter {
start_cursor,
user_ids,
paths: params.paths.into_iter().map(PathFilter::from).collect(),
mode: params.mode,
limit: params.limit,
})
}
pub async fn feed_stream(
State(state): State<AppState>,
RawQuery(raw_query): RawQuery,
) -> HttpResult<impl IntoResponse> {
let params = parse_admin_stream_query(raw_query.as_deref().unwrap_or(""))?;
let filter = resolve_filter(&state, params).await?;
let sse = Sse::new(
state
.events_service
.all_events_stream(state.sql_db.clone(), state.metrics.clone(), filter)
.map(|event| {
Ok::<_, Infallible>(
Event::default()
.event(event.event_type.to_string())
.data(event.to_sse_data()),
)
}),
)
.keep_alive(KeepAlive::default());
Ok((
[(header::CACHE_CONTROL, HeaderValue::from_static("no-store"))],
sse,
))
}