Skip to main content

max / makenotwork

7.0 KB · 222 lines History Blame Raw
1 //! Passkey / WebAuthn management API: register, list, rename, delete.
2
3 use axum::{
4 Form, Json,
5 extract::{Path, State},
6 response::{IntoResponse, Response},
7 };
8 use serde::Deserialize;
9 use tower_sessions::Session;
10 use webauthn_rs::prelude::*;
11
12 use sqlx::PgPool;
13
14 use crate::{
15 auth::{AuthUser, verify_password_async},
16 db::{self, PasskeyId},
17 error::{AppError, Result, ResultExt},
18 helpers::hx_toast,
19 templates::{PasskeyDisplay, PasskeyListTemplate},
20 };
21
22 /// Session key for in-flight passkey registration challenge state.
23 const PASSKEY_REG_STATE_KEY: &str = "passkey_reg_state";
24
25 /// Maximum number of passkeys a user can register.
26 const MAX_PASSKEYS_PER_USER: i64 = 20;
27
28 /// Form input for password confirmation on registration.
29 #[derive(Deserialize)]
30 pub(super) struct RegisterStartForm {
31 password: String,
32 }
33
34 /// Start passkey registration: generate challenge, return CreationChallengeResponse as JSON.
35 /// Requires password confirmation to prevent session-theft → persistent backdoor.
36 #[tracing::instrument(skip_all, name = "passkeys::register_start")]
37 pub(super) async fn register_start(
38 State(db): State<PgPool>,
39 State(webauthn): State<std::sync::Arc<Webauthn>>,
40 AuthUser(user): AuthUser,
41 session: Session,
42 Form(form): Form<RegisterStartForm>,
43 ) -> Result<Response> {
44 user.check_not_sandbox()?;
45
46 // Require password confirmation (matches delete flow)
47 let db_user = db::users::get_user_by_id(&db, user.id)
48 .await?
49 .ok_or(AppError::Unauthorized)?;
50 if !verify_password_async(form.password.clone(), db_user.password_hash.clone()).await? {
51 return Err(AppError::BadRequest("Incorrect password".to_string()));
52 }
53 // Enforce registration cap
54 let count = db::passkeys::count_passkeys(&db, user.id).await?;
55 if count >= MAX_PASSKEYS_PER_USER {
56 return Err(AppError::BadRequest(format!(
57 "Maximum of {MAX_PASSKEYS_PER_USER} passkeys reached"
58 )));
59 }
60
61 // Load existing credentials to exclude (prevents re-registration of the same authenticator)
62 let existing_json = db::passkeys::get_passkey_credentials(&db, user.id).await?;
63 let exclude_creds: Vec<CredentialID> = existing_json
64 .iter()
65 .filter_map(|j| serde_json::from_value::<Passkey>(j.clone()).ok())
66 .map(|pk| pk.cred_id().clone())
67 .collect();
68
69 let exclude = if exclude_creds.is_empty() {
70 None
71 } else {
72 Some(exclude_creds)
73 };
74
75 let (ccr, reg_state) = webauthn
76 .start_passkey_registration(
77 *user.id.as_uuid(),
78 user.username.as_ref(),
79 user.username.as_ref(),
80 exclude,
81 )
82 .context("webauthn registration start")?;
83
84 // Store the registration state in session for the finish step
85 session
86 .insert(PASSKEY_REG_STATE_KEY, &reg_state)
87 .await
88 .context("session error")?;
89
90 Ok(Json(ccr).into_response())
91 }
92
93 /// Finish passkey registration: verify attestation, store credential.
94 #[tracing::instrument(skip_all, name = "passkeys::register_finish")]
95 pub(super) async fn register_finish(
96 State(db): State<PgPool>,
97 State(webauthn): State<std::sync::Arc<Webauthn>>,
98 AuthUser(user): AuthUser,
99 session: Session,
100 Json(reg): Json<RegisterPublicKeyCredential>,
101 ) -> Result<Response> {
102 let reg_state: PasskeyRegistration = session
103 .get(PASSKEY_REG_STATE_KEY)
104 .await
105 .context("session error")?
106 .ok_or_else(|| AppError::BadRequest("No pending registration".to_string()))?;
107
108 // Clean up session state
109 session
110 .remove::<PasskeyRegistration>(PASSKEY_REG_STATE_KEY)
111 .await
112 .ok();
113
114 let passkey = webauthn
115 .finish_passkey_registration(&reg, &reg_state)
116 .map_err(|e| AppError::BadRequest(format!("Registration failed: {e}")))?;
117
118 let credential_json = serde_json::to_value(&passkey).context("serialize passkey")?;
119 let credential_id = passkey.cred_id().clone();
120
121 db::passkeys::create_passkey(&db, user.id, "Passkey", &credential_json, &credential_id)
122 .await
123 .map_err(|e| {
124 crate::helpers::map_unique_violation(e, "This passkey is already registered")
125 })?;
126
127 tracing::info!(user_id = %user.id, event = "passkey_registered", "Passkey registered");
128
129 Ok((
130 [("HX-Trigger", hx_toast("Passkey registered", "success"))],
131 list_inner(&db, user.id).await?,
132 )
133 .into_response())
134 }
135
136 /// List passkeys as HTMX partial.
137 #[tracing::instrument(skip_all, name = "passkeys::list")]
138 pub(super) async fn list(State(db): State<PgPool>, AuthUser(user): AuthUser) -> Result<Response> {
139 Ok(list_inner(&db, user.id).await?.into_response())
140 }
141
142 /// Inner helper to build the passkey list template.
143 async fn list_inner(db: &PgPool, user_id: db::UserId) -> Result<PasskeyListTemplate> {
144 let passkeys = db::passkeys::list_passkeys(db, user_id).await?;
145 let passkeys = passkeys
146 .into_iter()
147 .map(|p| PasskeyDisplay {
148 id: p.id.to_string(),
149 name: p.name,
150 created_at: p.created_at.format("%Y-%m-%d").to_string(),
151 last_used_at: p.last_used_at.map(|d| d.format("%Y-%m-%d").to_string()),
152 })
153 .collect();
154
155 Ok(PasskeyListTemplate { passkeys })
156 }
157
158 /// Rename a passkey.
159 #[derive(Deserialize)]
160 pub(super) struct RenameForm {
161 name: String,
162 }
163
164 #[tracing::instrument(skip_all, name = "passkeys::rename")]
165 pub(super) async fn rename(
166 State(db): State<PgPool>,
167 AuthUser(user): AuthUser,
168 Path(id): Path<PasskeyId>,
169 Form(form): Form<RenameForm>,
170 ) -> Result<Response> {
171 let name = form.name.trim();
172 if name.is_empty() || name.len() > 100 {
173 return Err(AppError::validation(
174 "Name must be 1-100 characters".to_string(),
175 ));
176 }
177
178 if !db::passkeys::rename_passkey(&db, id, user.id, name).await? {
179 return Err(AppError::NotFound);
180 }
181
182 Ok((
183 [("HX-Trigger", hx_toast("Passkey renamed", "success"))],
184 list_inner(&db, user.id).await?,
185 )
186 .into_response())
187 }
188
189 /// Delete a passkey (requires password confirmation).
190 #[derive(Deserialize)]
191 pub(super) struct DeleteForm {
192 password: String,
193 }
194
195 #[tracing::instrument(skip_all, name = "passkeys::delete")]
196 pub(super) async fn delete(
197 State(db): State<PgPool>,
198 AuthUser(user): AuthUser,
199 Path(id): Path<PasskeyId>,
200 Form(form): Form<DeleteForm>,
201 ) -> Result<Response> {
202 let db_user = db::users::get_user_by_id(&db, user.id)
203 .await?
204 .ok_or(AppError::Unauthorized)?;
205
206 if !verify_password_async(form.password.clone(), db_user.password_hash.clone()).await? {
207 return Err(AppError::BadRequest("Incorrect password".to_string()));
208 }
209
210 if !db::passkeys::delete_passkey(&db, id, user.id).await? {
211 return Err(AppError::NotFound);
212 }
213
214 tracing::info!(user_id = %user.id, passkey_id = %id, event = "passkey_deleted", "Passkey deleted");
215
216 Ok((
217 [("HX-Trigger", hx_toast("Passkey deleted", "success"))],
218 list_inner(&db, user.id).await?,
219 )
220 .into_response())
221 }
222