diff --git a/crates/server/src/http.rs b/crates/server/src/http.rs index c5f6097..9cb33ef 100644 --- a/crates/server/src/http.rs +++ b/crates/server/src/http.rs @@ -207,6 +207,22 @@ struct CreateStreamKeyBody { label: String, } +struct SessionCookie(Option); + +impl FromRequestParts for SessionCookie +where + S: Send + Sync, +{ + type Rejection = HttpError; + + async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result { + let jar = CookieJar::from_request_parts(parts, _state).await.unwrap(); + Ok(SessionCookie( + jar.get("session").map(|c| c.value().to_string()), + )) + } +} + struct AuthUser(entity::users::Model); impl FromRequestParts> for AuthUser { @@ -216,12 +232,18 @@ impl FromRequestParts> for AuthUser { parts: &mut Parts, state: &Arc, ) -> Result { - let jar = CookieJar::from_request_parts(parts, state).await.unwrap(); + let jar = CookieJar::from_headers(&parts.headers); let session = jar.get("session").ok_or(HttpError::Unauthorized)?; + let session_val = session.value().to_string(); + tracing::debug!(?session_val, "AuthUser: extracted session cookie"); - let user = users::Entity::find_by_auth_session(&state.db, session.to_string()) - .await? - .ok_or(HttpError::Unauthorized)?; + let user = match users::Entity::find_by_auth_session(&state.db, session_val.clone()).await { + Ok(u) => u.ok_or(HttpError::Unauthorized)?, + Err(e) => { + tracing::error!(?session_val, error = %e, "AuthUser: db query failed"); + return Err(HttpError::DbErr(e)); + } + }; Ok(AuthUser(user)) } @@ -289,11 +311,45 @@ struct LoginResponse { session_token: String, } +struct DevFlag(bool); + +impl FromRequestParts for DevFlag +where + S: Send + Sync, +{ + type Rejection = std::convert::Infallible; + + async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result { + let query = parts.uri.query().unwrap_or(""); + Ok(DevFlag(query.contains("dev=1"))) + } +} + +fn cookie_for_token(token: &str, dev: bool) -> Cookie<'static> { + let token = token.to_owned(); + match dev { + true => Cookie::build(("session", token)) + .path("/") + .http_only(true) + .same_site(axum_extra::extract::cookie::SameSite::None) + .build(), + false => Cookie::build(("session", token)) + .path("/") + .http_only(true) + .same_site(axum_extra::extract::cookie::SameSite::Lax) + .build(), + } +} + +#[axum::debug_handler] async fn login_handler( + session: SessionCookie, State(state): State>, + DevFlag(dev): DevFlag, Json(payload): Json, - jar: CookieJar, ) -> Result { + tracing::debug!(session = ?session.0, "login: existing session cookie"); + let user = users::Entity::find_by_username(&state.db, payload.username.clone()) .await? .ok_or_else(|| { @@ -310,8 +366,11 @@ async fn login_handler( let auth = auth_session::Entity::create(&state.db, user.id).await?; let token = auth.value; - let jar = jar.add(Cookie::new("session", token)); - Ok((StatusCode::OK, jar)) + let cookie = cookie_for_token(&token, dev); + Ok(( + StatusCode::OK, + [(axum::http::header::SET_COOKIE, cookie.to_string())], + )) } #[derive(Deserialize)] @@ -322,10 +381,12 @@ struct CreateUserForm { } async fn create_user_handler( + session: SessionCookie, State(state): State>, + DevFlag(dev): DevFlag, Json(payload): Json, - jar: CookieJar, ) -> Result { + tracing::debug!(session = ?session.0, "create_user: existing session cookie"); if state.config.signup_code.is_empty() || payload.ref_token != state.config.signup_code { warn!(username = %payload.username, "signup rejected: invalid signup code"); return Err(HttpError::Unauthorized); @@ -352,8 +413,11 @@ async fn create_user_handler( let session = auth_session::Entity::create(&state.db, user.id).await?; let token = session.value; - let jar = jar.add(Cookie::new("session", token)); - Ok((jar, StatusCode::OK)) + let cookie = cookie_for_token(&token, dev); + Ok(( + StatusCode::OK, + [(axum::http::header::SET_COOKIE, cookie.to_string())], + )) } async fn stream_handler( @@ -372,7 +436,7 @@ async fn stream_handler( }; let stream_key_id = stream_key_id.ok_or_else(|| { - warn!(slug = %slug, "WHEP request for unknown or inactive stream"); + // warn!(slug = %slug, "WHEP request for unknown or inactive stream"); HttpError::NotFound })?;