a lot 2
This commit is contained in:
+150
-48
@@ -4,24 +4,35 @@ use axum::{
|
||||
Json, Router,
|
||||
body::{Body, Bytes},
|
||||
extract::{Form, FromRequestParts, Path, Query, State},
|
||||
http::{HeaderMap, StatusCode, header::SET_COOKIE, request::Parts},
|
||||
http::{
|
||||
HeaderMap, HeaderName, Method, StatusCode,
|
||||
header::{ACCEPT, AUTHORIZATION, CONTENT_TYPE, SET_COOKIE},
|
||||
request::Parts,
|
||||
},
|
||||
response::{IntoResponse, Response},
|
||||
routing::{get, post},
|
||||
};
|
||||
use entity::{auth_session, stream_key, users};
|
||||
use sea_orm::DatabaseConnection;
|
||||
use entity::{auth_session, stream_key, stream_session, users};
|
||||
use sea_orm::{DatabaseConnection, EntityTrait, QueryFilter};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::{
|
||||
net::TcpListener,
|
||||
sync::{Mutex, mpsc::Sender},
|
||||
};
|
||||
use tower_http::cors::CorsLayer;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::{AppState, hash::hash_password};
|
||||
use tracing::{debug, info};
|
||||
|
||||
use crate::{
|
||||
AppState,
|
||||
hash::{hash_password, verify_password},
|
||||
webrtc_ingest::handle_whip_injest,
|
||||
};
|
||||
|
||||
pub struct HttpServer {
|
||||
pub offer_tx: Sender<(i32, i32, String)>,
|
||||
pub accept_rx: async_broadcast::Receiver<(i32, Option<String>)>,
|
||||
pub accept_rx: async_broadcast::InactiveReceiver<(i32, Option<String>)>,
|
||||
pub appstate: Arc<Mutex<AppState>>,
|
||||
pub request_count: Mutex<i32>,
|
||||
pub db: DatabaseConnection,
|
||||
@@ -31,6 +42,22 @@ impl HttpServer {
|
||||
pub fn start(self) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let state = Arc::new(self);
|
||||
|
||||
let origins = [
|
||||
"http://localhost:5173".parse().unwrap(),
|
||||
"https://stream.h.doloro.co.uk".parse().unwrap(),
|
||||
];
|
||||
|
||||
let cors = CorsLayer::new()
|
||||
.allow_methods([Method::GET, Method::POST, Method::OPTIONS])
|
||||
.allow_headers([
|
||||
AUTHORIZATION,
|
||||
ACCEPT,
|
||||
CONTENT_TYPE,
|
||||
HeaderName::from_static("session"),
|
||||
])
|
||||
.allow_credentials(true)
|
||||
.allow_origin(origins);
|
||||
|
||||
tokio::spawn(async move {
|
||||
let app = Router::new()
|
||||
.route("/api/catalog", get(catalog_handler))
|
||||
@@ -39,9 +66,11 @@ impl HttpServer {
|
||||
"/api/stream-key",
|
||||
post(create_stream_key_handler).get(get_all_stream_keys),
|
||||
)
|
||||
.route("/api/whip", post(handle_whip_injest))
|
||||
.route("/api/login", post(login_handler))
|
||||
.route("/api/stream/{slug}", post(stream_handler))
|
||||
.route("/api/meow", get(meow_handler))
|
||||
.layer(cors)
|
||||
.with_state(state);
|
||||
|
||||
let addr = SocketAddr::from(([127, 0, 0, 1], 3000));
|
||||
@@ -52,8 +81,45 @@ impl HttpServer {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct StreamCatalog {
|
||||
active_streams: Vec<StreamListing>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct StreamListing {
|
||||
label: String,
|
||||
id: i32,
|
||||
user: String,
|
||||
}
|
||||
|
||||
async fn catalog_handler(State(state): State<Arc<HttpServer>>) -> impl IntoResponse {
|
||||
let catalog = catalog_from_state(&state.appstate).await;
|
||||
let streams = stream_session::Model::get_all_active_sessions(&state.db)
|
||||
.await
|
||||
.unwrap();
|
||||
debug!("{:#?}", streams);
|
||||
let catalog: Vec<StreamListing> = futures::future::join_all(streams.iter().map(|listing| {
|
||||
let db = state.db.clone();
|
||||
let stream_key_id = listing.stream_key_id;
|
||||
async move {
|
||||
let key_info = stream_key::Entity::find_by_id(stream_key_id)
|
||||
.one(&db)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let user = users::Entity::find_by_id(key_info.user_id)
|
||||
.one(&db)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
StreamListing {
|
||||
id: stream_key_id,
|
||||
label: key_info.label,
|
||||
user: user.username,
|
||||
}
|
||||
}
|
||||
}))
|
||||
.await;
|
||||
Json(catalog)
|
||||
}
|
||||
|
||||
@@ -62,6 +128,22 @@ struct CreateStreamKeyBody {
|
||||
label: String,
|
||||
}
|
||||
|
||||
fn extract_session_token(headers: &HeaderMap) -> Option<String> {
|
||||
// Check bare `session` header first (curl / API clients)
|
||||
if let Some(v) = headers.get("session") {
|
||||
return Some(v.to_str().ok()?.to_string());
|
||||
}
|
||||
// Fall back to Cookie header (browsers)
|
||||
let cookie_header = headers.get("cookie")?.to_str().ok()?;
|
||||
cookie_header
|
||||
.split(';')
|
||||
.find_map(|pair| {
|
||||
let pair = pair.trim();
|
||||
pair.strip_prefix("session=")
|
||||
})
|
||||
.map(|v| v.to_string())
|
||||
}
|
||||
|
||||
struct AuthUser(entity::users::Model);
|
||||
|
||||
impl FromRequestParts<Arc<HttpServer>> for AuthUser {
|
||||
@@ -71,17 +153,12 @@ impl FromRequestParts<Arc<HttpServer>> for AuthUser {
|
||||
parts: &mut Parts,
|
||||
state: &Arc<HttpServer>,
|
||||
) -> Result<Self, StatusCode> {
|
||||
let token = parts
|
||||
.headers
|
||||
.get("session")
|
||||
// parse token...
|
||||
.ok_or(StatusCode::UNAUTHORIZED)?;
|
||||
let token = extract_session_token(&parts.headers).ok_or(StatusCode::UNAUTHORIZED)?;
|
||||
|
||||
let user =
|
||||
users::Entity::find_by_auth_session(&state.db, token.to_str().unwrap().to_string())
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
|
||||
.ok_or(StatusCode::UNAUTHORIZED)?;
|
||||
let user = users::Entity::find_by_auth_session(&state.db, token)
|
||||
.await
|
||||
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
|
||||
.ok_or(StatusCode::UNAUTHORIZED)?;
|
||||
|
||||
Ok(AuthUser(user))
|
||||
}
|
||||
@@ -94,6 +171,13 @@ async fn create_stream_key_handler(
|
||||
) -> impl IntoResponse {
|
||||
let uuid = uuid::Uuid::new_v4();
|
||||
let value = format!("stream-key-{uuid}");
|
||||
let key_amount = stream_key::Entity::find_by_user(&state.db, auth.0.id)
|
||||
.await
|
||||
.unwrap()
|
||||
.len();
|
||||
if key_amount >= auth.0.stream_key_limit.try_into().unwrap() {
|
||||
return StatusCode::NOT_ACCEPTABLE;
|
||||
}
|
||||
let key = stream_key::Entity::create(&state.db, auth.0.id, value, payload.label, false).await;
|
||||
|
||||
if let Ok(_key) = key {
|
||||
@@ -117,9 +201,9 @@ async fn get_all_stream_keys(
|
||||
State(state): State<Arc<HttpServer>>,
|
||||
auth: AuthUser,
|
||||
) -> impl IntoResponse {
|
||||
let user_keys = stream_key::Entity::find_by_user(&state.db, auth.0.id)
|
||||
.await
|
||||
.unwrap();
|
||||
// let user_keys = stream_key::Entity::find_by_user(&state.db, auth.0.id)
|
||||
// .await
|
||||
// .unwrap();
|
||||
|
||||
match stream_key::Entity::find_by_user(&state.db, auth.0.id).await {
|
||||
Ok(keys) => Json(keys).into_response(),
|
||||
@@ -141,12 +225,36 @@ struct LoginResponse {
|
||||
async fn login_handler(
|
||||
State(state): State<Arc<HttpServer>>,
|
||||
Json(payload): Json<LoginForm>,
|
||||
) -> StatusCode {
|
||||
// TODO: hash payload.password with a real password hasher (argon2/bcrypt), look up user by username
|
||||
// TODO: verify hash matches stored hashed_password
|
||||
// TODO: call auth_session::Entity::create, return session token
|
||||
// TODO: return 401 on bad credentials
|
||||
todo!("login not yet implemented")
|
||||
) -> impl IntoResponse {
|
||||
if let Ok(x) = users::Entity::find_by_username(&state.db, payload.username).await {
|
||||
if let Some(x) = x {
|
||||
let pass = verify_password(&payload.password, &x.hashed_password);
|
||||
if !pass {
|
||||
let mut meow = Response::new("".to_string());
|
||||
*meow.status_mut() = StatusCode::UNAUTHORIZED;
|
||||
return meow;
|
||||
};
|
||||
let auth = auth_session::Entity::create(&state.db, x.id).await.unwrap(); // This should be ok (hopefully)
|
||||
let token = auth.value;
|
||||
let mut meow = Response::new("".to_string());
|
||||
meow.headers_mut().insert(
|
||||
SET_COOKIE,
|
||||
format!("session={token}; HttpOnly; SameSite=Strict; Path=/")
|
||||
.parse()
|
||||
.unwrap(),
|
||||
);
|
||||
*meow.status_mut() = StatusCode::OK;
|
||||
return meow;
|
||||
} else {
|
||||
let mut meow = Response::new("".to_string());
|
||||
*meow.status_mut() = StatusCode::UNAUTHORIZED;
|
||||
return meow;
|
||||
}
|
||||
} else {
|
||||
let mut meow = Response::new("".to_string());
|
||||
*meow.status_mut() = StatusCode::UNAUTHORIZED;
|
||||
return meow;
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
@@ -204,7 +312,7 @@ async fn stream_handler(
|
||||
let app = state.appstate.lock().await;
|
||||
app.stream_sessions
|
||||
.iter()
|
||||
.find(|e| e.value().stream_key_label == slug)
|
||||
.find(|e| e.value().stream_key_id.to_string() == slug)
|
||||
.map(|e| *e.key())
|
||||
};
|
||||
|
||||
@@ -218,31 +326,41 @@ async fn stream_handler(
|
||||
.unwrap();
|
||||
};
|
||||
|
||||
let mut accept_rx = state.accept_rx.new_receiver();
|
||||
let mut accept_rx = state.accept_rx.activate_cloned();
|
||||
let _ = state
|
||||
.offer_tx
|
||||
.send((request_id_clone, stream_key_id, body))
|
||||
.await
|
||||
.unwrap();
|
||||
debug!(
|
||||
request_id = request_id_clone,
|
||||
"offer sent, waiting for answer"
|
||||
);
|
||||
|
||||
let mut reply_body = String::new();
|
||||
while let Ok(answer) = accept_rx.recv().await {
|
||||
debug!(request_id = request_id_clone, "received answer candidate");
|
||||
if let Some(reply) = answer.1 {
|
||||
if answer.0 == request_id_clone {
|
||||
Response::builder()
|
||||
return Response::builder()
|
||||
.status(StatusCode::CREATED)
|
||||
.header("content-type", "application/sdp")
|
||||
.body(reply)
|
||||
.unwrap();
|
||||
} else {
|
||||
continue;
|
||||
}
|
||||
} else {
|
||||
Response::builder()
|
||||
} else if answer.0 == request_id_clone {
|
||||
return Response::builder()
|
||||
.status(StatusCode::NOT_FOUND)
|
||||
// .header("content-type", "application/sdp")
|
||||
.body("")
|
||||
.body(String::new())
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
info!(
|
||||
request_id = request_id_clone,
|
||||
"answer channel closed without a match"
|
||||
);
|
||||
|
||||
Response::builder()
|
||||
.status(StatusCode::CREATED)
|
||||
@@ -254,19 +372,3 @@ async fn stream_handler(
|
||||
async fn meow_handler() -> &'static str {
|
||||
"meow"
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct StreamCatalog {
|
||||
active_streams: Vec<i32>,
|
||||
}
|
||||
|
||||
// TODO: Make this return streak labels and their id instead of the whole stream key omfg
|
||||
async fn catalog_from_state(state: &Arc<Mutex<AppState>>) -> StreamCatalog {
|
||||
let app = state.lock().await;
|
||||
let active_streams = app
|
||||
.stream_sessions
|
||||
.iter()
|
||||
.map(|e| e.key().clone())
|
||||
.collect();
|
||||
StreamCatalog { active_streams }
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user