| 1 | use crate::*; |
| 2 | use base64::{ |
| 3 | Engine, |
| 4 | engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}, |
| 5 | }; |
| 6 | use openssl::{hash::MessageDigest, pkey::PKey, rsa::Rsa, sign::Signer}; |
| 7 | use rusqlite::{Connection, OptionalExtension, params as sql}; |
| 8 | use sha2::{Digest, Sha256}; |
| 9 | |
| 10 | const SCOPES: &[&str] = &["openid", "profile", "email", "groups", "offline_access"]; |
| 11 | const ACCESS_TTL: i64 = 300; |
| 12 | |
| 13 | pub fn initialise(db: &Connection) -> Result<()> { |
| 14 | db.execute_batch("CREATE TABLE IF NOT EXISTS oidc_clients (id TEXT PRIMARY KEY, secret_hash TEXT NOT NULL, config TEXT NOT NULL); |
| 15 | CREATE TABLE IF NOT EXISTS oidc_key (id INTEGER PRIMARY KEY CHECK(id=1), pem TEXT NOT NULL); |
| 16 | CREATE TABLE IF NOT EXISTS oidc_tokens (hash TEXT PRIMARY KEY,kind TEXT NOT NULL,user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,client_id TEXT NOT NULL REFERENCES oidc_clients(id) ON DELETE CASCADE,session_hash TEXT NOT NULL REFERENCES sessions(hash) ON DELETE CASCADE,expires INTEGER NOT NULL,family TEXT NOT NULL,scope TEXT NOT NULL,auth_time INTEGER NOT NULL); |
| 17 | CREATE INDEX IF NOT EXISTS oidc_token_family ON oidc_tokens(family);")?; |
| 18 | let exists: bool = db.query_row("SELECT EXISTS(SELECT 1 FROM oidc_key)", [], |r| r.get(0))?; |
| 19 | if !exists { |
| 20 | let pem = PKey::from_rsa(Rsa::generate(2048)?)?.private_key_to_pem_pkcs8()?; |
| 21 | db.execute( |
| 22 | "INSERT OR IGNORE INTO oidc_key VALUES (1,?)", |
| 23 | [String::from_utf8(pem)?], |
| 24 | )?; |
| 25 | } |
| 26 | Ok(()) |
| 27 | } |
| 28 | |
| 29 | fn invalid(code: &str) -> Error { |
| 30 | Error::new(400, code) |
| 31 | } |
| 32 | fn fields(value: &str) -> Result<HashMap<String, String>> { |
| 33 | let mut fields = HashMap::new(); |
| 34 | for (name, value) in url::form_urlencoded::parse(value.as_bytes()) { |
| 35 | if fields |
| 36 | .insert(name.into_owned(), value.into_owned()) |
| 37 | .is_some() |
| 38 | { |
| 39 | return Err(invalid("invalid_request")); |
| 40 | } |
| 41 | } |
| 42 | Ok(fields) |
| 43 | } |
| 44 | fn client(db: &Connection, id: &str) -> Result<Value> { |
| 45 | let value: Option<String> = db |
| 46 | .query_row("SELECT config FROM oidc_clients WHERE id=?", [id], |r| { |
| 47 | r.get(0) |
| 48 | }) |
| 49 | .optional()?; |
| 50 | value |
| 51 | .map(|s| serde_json::from_str(&s).map_err(Error::from)) |
| 52 | .transpose()? |
| 53 | .ok_or_else(|| invalid("invalid_client")) |
| 54 | } |
| 55 | |
| 56 | fn guests_allowed(id: &str, config: &Value) -> bool { |
| 57 | config["allowGuests"] == true && (id == "shale" || id.starts_with("shale-preview-")) |
| 58 | } |
| 59 | |
| 60 | pub(crate) fn sign_in_target(auth: &auth::Store, next: &str) -> Result<(url::Url, Value)> { |
| 61 | if next.len() > 8192 || next.contains('\\') || next.chars().any(char::is_control) { |
| 62 | return Err(invalid("invalid_request")); |
| 63 | } |
| 64 | let target = auth |
| 65 | .origin |
| 66 | .join(next) |
| 67 | .map_err(|_| invalid("invalid_request"))?; |
| 68 | if target.origin() != auth.origin.origin() |
| 69 | || target.path() != "/auth/oidc/authorize" |
| 70 | || target.fragment().is_some() |
| 71 | || !target.username().is_empty() |
| 72 | || target.password().is_some() |
| 73 | { |
| 74 | return Err(invalid("invalid_request")); |
| 75 | } |
| 76 | let query = fields(target.query().unwrap_or_default())?; |
| 77 | let id = query |
| 78 | .get("client_id") |
| 79 | .map(String::as_str) |
| 80 | .unwrap_or_default(); |
| 81 | let config = client(&auth.db.lock().unwrap(), id)?; |
| 82 | if query.get("response_type").map(String::as_str) != Some("code") |
| 83 | || !array(&config["redirectUris"]) |
| 84 | .iter() |
| 85 | .any(|v| v.as_str() == query.get("redirect_uri").map(String::as_str)) |
| 86 | { |
| 87 | return Err(invalid("access_denied")); |
| 88 | } |
| 89 | Ok((target, config)) |
| 90 | } |
| 91 | |
| 92 | pub(crate) fn guest_target(auth: &auth::Store, next: &str) -> Result<String> { |
| 93 | let (target, config) = sign_in_target(auth, next)?; |
| 94 | let query = fields(target.query().unwrap_or_default())?; |
| 95 | if !guests_allowed( |
| 96 | query |
| 97 | .get("client_id") |
| 98 | .map(String::as_str) |
| 99 | .unwrap_or_default(), |
| 100 | &config, |
| 101 | ) { |
| 102 | return Err(invalid("access_denied")); |
| 103 | } |
| 104 | Ok(format!( |
| 105 | "{}?{}", |
| 106 | target.path(), |
| 107 | target.query().unwrap_or_default() |
| 108 | )) |
| 109 | } |
| 110 | |
| 111 | pub fn username(auth: &auth::Store, user_id: &str, client_id: &str) -> Result<String> { |
| 112 | let db = auth.db.lock().unwrap(); |
| 113 | let user = auth::user(&db, user_id)?; |
| 114 | if user["enabled"] != true { |
| 115 | return Err(Error::new(403, "This account is disabled.")); |
| 116 | } |
| 117 | let config = client(&db, client_id)?; |
| 118 | let username = string(&user["username"]); |
| 119 | Ok(config["usernameAliases"][username] |
| 120 | .as_str() |
| 121 | .unwrap_or(username) |
| 122 | .to_owned()) |
| 123 | } |
| 124 | |
| 125 | /// Only the root-owned deployment CLI can register first-party clients. |
| 126 | pub fn provision(auth: &auth::Store, input: Value) -> Result<Value> { |
| 127 | let request = &input["request"]; |
| 128 | let id = string(&request["clientId"]); |
| 129 | let stage = string(&input["stageId"]); |
| 130 | if id.is_empty() || id.len() > 128 || !id.chars().all(|c| c.is_ascii_alphanumeric() || c == '-') |
| 131 | { |
| 132 | return Err(invalid("invalid_client")); |
| 133 | } |
| 134 | let mut db = auth.db.lock().unwrap(); |
| 135 | let tx = db.transaction()?; |
| 136 | if input["operation"] == "delete" { |
| 137 | if stage.is_empty() || id != stage { |
| 138 | return Err(invalid("invalid_client")); |
| 139 | } |
| 140 | tx.execute("DELETE FROM oidc_clients WHERE id=?", [id])?; |
| 141 | tx.commit()?; |
| 142 | return Ok(json!({})); |
| 143 | } |
| 144 | let redirects = array(&request["redirectUris"]); |
| 145 | if redirects.is_empty() || redirects.len() > 64 { |
| 146 | return Err(invalid("invalid_redirect_uri")); |
| 147 | } |
| 148 | for value in redirects { |
| 149 | let value = value |
| 150 | .as_str() |
| 151 | .ok_or_else(|| invalid("invalid_redirect_uri"))?; |
| 152 | let uri = url::Url::parse(value).map_err(|_| invalid("invalid_redirect_uri"))?; |
| 153 | let site = auth |
| 154 | .origin |
| 155 | .host_str() |
| 156 | .unwrap() |
| 157 | .strip_prefix("snowglobe.") |
| 158 | .unwrap_or_default(); |
| 159 | if uri.scheme() != "https" |
| 160 | || site.is_empty() |
| 161 | || !uri |
| 162 | .host_str() |
| 163 | .is_some_and(|h| h.ends_with(&format!(".{site}"))) |
| 164 | || !uri.username().is_empty() |
| 165 | || uri.password().is_some() |
| 166 | || uri.fragment().is_some() |
| 167 | || value.contains('*') |
| 168 | { |
| 169 | return Err(invalid("invalid_redirect_uri")); |
| 170 | } |
| 171 | } |
| 172 | if !stage.is_empty() |
| 173 | && (id != stage |
| 174 | || redirects.iter().any(|v| { |
| 175 | url::Url::parse(string(v)) |
| 176 | .ok() |
| 177 | .and_then(|u| u.host_str().map(str::to_owned)) |
| 178 | .is_none_or(|h| !h.starts_with(&format!("{stage}."))) |
| 179 | })) |
| 180 | { |
| 181 | return Err(invalid("invalid_redirect_uri")); |
| 182 | } |
| 183 | let aliases = request["usernameAliases"] |
| 184 | .as_object() |
| 185 | .cloned() |
| 186 | .unwrap_or_default(); |
| 187 | let mut unique = std::collections::HashSet::new(); |
| 188 | for (name, alias) in &aliases { |
| 189 | let alias = alias |
| 190 | .as_str() |
| 191 | .filter(|a| !a.is_empty()) |
| 192 | .ok_or_else(|| invalid("invalid_alias"))?; |
| 193 | let matches: i64 = |
| 194 | tx.query_row("SELECT count(*) FROM users WHERE username=?", [name], |r| { |
| 195 | r.get(0) |
| 196 | })?; |
| 197 | if matches != 1 || !unique.insert(alias) { |
| 198 | return Err(invalid("invalid_alias")); |
| 199 | } |
| 200 | } |
| 201 | let previous = &input["existing"]; |
| 202 | if previous["clientId"].as_str().is_some_and(|old| old != id) { |
| 203 | return Err(invalid("invalid_client")); |
| 204 | } |
| 205 | let secret = previous["clientSecret"] |
| 206 | .as_str() |
| 207 | .filter(|s| s.len() >= 24) |
| 208 | .map(str::to_owned) |
| 209 | .unwrap_or_else(mcp::secret); |
| 210 | let allow_guests = request["allowGuests"] == true; |
| 211 | if allow_guests |
| 212 | && (!(id == "shale" || id.starts_with("shale-preview-")) |
| 213 | || redirects.iter().any(|v| { |
| 214 | url::Url::parse(string(v)).is_ok_and(|u| { |
| 215 | let site = auth |
| 216 | .origin |
| 217 | .host_str() |
| 218 | .unwrap() |
| 219 | .strip_prefix("snowglobe.") |
| 220 | .unwrap_or_default(); |
| 221 | u.path() != "/-/callback" |
| 222 | || u.host_str() != Some(format!("{id}.{site}").as_str()) |
| 223 | || u.query().is_some() |
| 224 | }) |
| 225 | })) |
| 226 | { |
| 227 | return Err(invalid("access_denied")); |
| 228 | } |
| 229 | let config = json!({"name":request["name"], "redirectUris":redirects, "usernameAliases":aliases, "allowGuests":allow_guests, "usernameRequired":request["usernameRequired"]==true}); |
| 230 | let old: Option<(String, String)> = tx |
| 231 | .query_row( |
| 232 | "SELECT secret_hash,config FROM oidc_clients WHERE id=?", |
| 233 | [id], |
| 234 | |r| Ok((r.get(0)?, r.get(1)?)), |
| 235 | ) |
| 236 | .optional()?; |
| 237 | if old |
| 238 | .as_ref() |
| 239 | .is_some_and(|old| old != &(mcp::hash(&secret), config.to_string())) |
| 240 | { |
| 241 | tx.execute("DELETE FROM oidc_tokens WHERE client_id=?", [id])?; |
| 242 | tx.execute( |
| 243 | "DELETE FROM pending WHERE kind='oidc-code' AND json_extract(data,'$.client')=?", |
| 244 | [id], |
| 245 | )?; |
| 246 | } |
| 247 | tx.execute("INSERT INTO oidc_clients VALUES (?,?,?) ON CONFLICT(id) DO UPDATE SET secret_hash=excluded.secret_hash,config=excluded.config", sql![id,mcp::hash(&secret),config.to_string()])?; |
| 248 | tx.commit()?; |
| 249 | Ok( |
| 250 | json!({"clientId":id,"clientSecret":secret,"issuerUrl":auth.origin.origin().ascii_serialization()}), |
| 251 | ) |
| 252 | } |
| 253 | |
| 254 | fn public_key(db: &Connection) -> Result<(PKey<openssl::pkey::Private>, String)> { |
| 255 | let pem: String = db.query_row("SELECT pem FROM oidc_key WHERE id=1", [], |r| r.get(0))?; |
| 256 | let key = PKey::private_key_from_pem(pem.as_bytes())?; |
| 257 | let kid = URL_SAFE_NO_PAD.encode(Sha256::digest(key.public_key_to_der()?)); |
| 258 | Ok((key, kid)) |
| 259 | } |
| 260 | fn jwt(db: &Connection, claims: &Value) -> Result<String> { |
| 261 | let (key, kid) = public_key(db)?; |
| 262 | let data = format!( |
| 263 | "{}.{}", |
| 264 | URL_SAFE_NO_PAD.encode(serde_json::to_vec( |
| 265 | &json!({"alg":"RS256","typ":"JWT","kid":kid}) |
| 266 | )?), |
| 267 | URL_SAFE_NO_PAD.encode(serde_json::to_vec(claims)?) |
| 268 | ); |
| 269 | let mut signer = Signer::new(MessageDigest::sha256(), &key)?; |
| 270 | signer.update(data.as_bytes())?; |
| 271 | Ok(format!( |
| 272 | "{data}.{}", |
| 273 | URL_SAFE_NO_PAD.encode(signer.sign_to_vec()?) |
| 274 | )) |
| 275 | } |
| 276 | fn claims(user: &Value, config: &Value, scopes: &str) -> Value { |
| 277 | let mut value = json!({"sub":user["id"]}); |
| 278 | let scopes: Vec<_> = scopes.split_whitespace().collect(); |
| 279 | // Shale requests only openid but needs a stable username to bind its account. |
| 280 | if scopes.contains(&"profile") || config["usernameRequired"] == true { |
| 281 | let username = string(&user["username"]); |
| 282 | value["preferred_username"] = config["usernameAliases"][username] |
| 283 | .as_str() |
| 284 | .map(|v| json!(v)) |
| 285 | .unwrap_or_else(|| json!(username)); |
| 286 | } |
| 287 | if scopes.contains(&"profile") { |
| 288 | let name = format!( |
| 289 | "{} {}", |
| 290 | string(&user["firstName"]), |
| 291 | string(&user["lastName"]) |
| 292 | ); |
| 293 | if !name.trim().is_empty() { |
| 294 | value["name"] = json!(name.trim()); |
| 295 | } |
| 296 | for (claim, field) in [("given_name", "firstName"), ("family_name", "lastName")] { |
| 297 | if !string(&user[field]).is_empty() { |
| 298 | value[claim] = user[field].clone(); |
| 299 | } |
| 300 | } |
| 301 | } |
| 302 | if scopes.contains(&"email") && !string(&user["email"]).is_empty() { |
| 303 | value["email"] = user["email"].clone(); |
| 304 | value["email_verified"] = json!(user["emailVerified"] == true); |
| 305 | } |
| 306 | if scopes.contains(&"groups") { |
| 307 | value["groups"] = json!( |
| 308 | array(&user["groups"]) |
| 309 | .iter() |
| 310 | .map(|g| format!("role:{}", string(&g["name"]))) |
| 311 | .collect::<Vec<_>>() |
| 312 | ); |
| 313 | } |
| 314 | value |
| 315 | } |
| 316 | fn eligible(db: &Connection, user_id: &str, session: &str) -> Result<Value> { |
| 317 | let valid: bool = db.query_row("SELECT EXISTS(SELECT 1 FROM sessions WHERE hash=? AND user_id=? AND client='dashboard' AND expires>?)",sql![session,user_id,now() as i64],|r|r.get(0))?; |
| 318 | let user = auth::user(db, user_id).map_err(|_| invalid("invalid_grant"))?; |
| 319 | if !valid || user["enabled"] != true || !array(&user["requiredActions"]).is_empty() { |
| 320 | return Err(invalid("invalid_grant")); |
| 321 | } |
| 322 | Ok(user) |
| 323 | } |
| 324 | fn authenticated_client( |
| 325 | db: &Connection, |
| 326 | headers: &HeaderMap, |
| 327 | form: &HashMap<String, String>, |
| 328 | ) -> Result<String> { |
| 329 | let (id, secret) = if let Some(header) = headers.get("authorization") { |
| 330 | let header = header |
| 331 | .to_str() |
| 332 | .ok() |
| 333 | .and_then(|h| h.strip_prefix("Basic ")) |
| 334 | .ok_or_else(|| invalid("invalid_client"))?; |
| 335 | let raw = String::from_utf8( |
| 336 | STANDARD |
| 337 | .decode(header) |
| 338 | .map_err(|_| invalid("invalid_client"))?, |
| 339 | ) |
| 340 | .map_err(|_| invalid("invalid_client"))?; |
| 341 | let (id, secret) = raw |
| 342 | .split_once(':') |
| 343 | .ok_or_else(|| invalid("invalid_client"))?; |
| 344 | let decode = |s: &str| -> String { |
| 345 | url::form_urlencoded::parse(format!("x={s}").as_bytes()) |
| 346 | .next() |
| 347 | .unwrap() |
| 348 | .1 |
| 349 | .into_owned() |
| 350 | }; |
| 351 | let (id, secret) = (decode(id), decode(secret)); |
| 352 | // Shale repeats Basic credentials in its form; require the same secret. |
| 353 | if form.get("client_secret").is_some_and(|supplied| { |
| 354 | !(id == "shale" || id.starts_with("shale-preview-")) |
| 355 | || mcp::hash(supplied) != mcp::hash(&secret) |
| 356 | }) { |
| 357 | return Err(invalid("invalid_request")); |
| 358 | } |
| 359 | (id, secret) |
| 360 | } else { |
| 361 | ( |
| 362 | form.get("client_id").cloned().unwrap_or_default(), |
| 363 | form.get("client_secret").cloned().unwrap_or_default(), |
| 364 | ) |
| 365 | }; |
| 366 | if secret.is_empty() |
| 367 | || form |
| 368 | .get("client_id") |
| 369 | .is_some_and(|supplied| supplied != &id) |
| 370 | { |
| 371 | return Err(invalid("invalid_client")); |
| 372 | } |
| 373 | let stored: Option<String> = db |
| 374 | .query_row( |
| 375 | "SELECT secret_hash FROM oidc_clients WHERE id=?", |
| 376 | [&id], |
| 377 | |r| r.get(0), |
| 378 | ) |
| 379 | .optional()?; |
| 380 | if stored.is_none_or(|s| !bool::from(s.as_bytes().ct_eq(mcp::hash(&secret).as_bytes()))) { |
| 381 | return Err(invalid("invalid_client")); |
| 382 | } |
| 383 | Ok(id) |
| 384 | } |
| 385 | fn tokens(db: &Connection, auth: &auth::Store, data: &Value, nonce: Option<&str>) -> Result<Value> { |
| 386 | let user = eligible(db, string(&data["user"]), string(&data["session"]))?; |
| 387 | let config = client(db, string(&data["client"]))?; |
| 388 | if guest::is_guest(&user) && !guests_allowed(string(&data["client"]), &config) { |
| 389 | return Err(invalid("access_denied")); |
| 390 | } |
| 391 | let scope = string(&data["scope"]); |
| 392 | let access = mcp::secret(); |
| 393 | let family = string(&data["family"]); |
| 394 | let issued = now() as i64; |
| 395 | let mut value = claims(&user, &config, scope); |
| 396 | value["iss"] = json!(auth.origin.origin().ascii_serialization()); |
| 397 | value["aud"] = data["client"].clone(); |
| 398 | value["iat"] = json!(issued); |
| 399 | value["exp"] = json!(issued + ACCESS_TTL); |
| 400 | value["auth_time"] = data["auth_time"].clone(); |
| 401 | value["at_hash"] = json!(URL_SAFE_NO_PAD.encode(&Sha256::digest(access.as_bytes())[..16])); |
| 402 | if let Some(nonce) = nonce { |
| 403 | value["nonce"] = json!(nonce); |
| 404 | } |
| 405 | let id_token = jwt(db, &value)?; |
| 406 | let mut output = json!({"access_token":access,"token_type":"Bearer","expires_in":ACCESS_TTL,"id_token":id_token,"scope":scope}); |
| 407 | for (kind, token, expires) in [ |
| 408 | ("access", access, issued + ACCESS_TTL), |
| 409 | ("refresh", mcp::secret(), issued + 30 * 86400), |
| 410 | ] { |
| 411 | db.execute( |
| 412 | "INSERT INTO oidc_tokens VALUES (?,?,?,?,?,?,?,?,?)", |
| 413 | sql![ |
| 414 | mcp::hash(&token), |
| 415 | kind, |
| 416 | string(&data["user"]), |
| 417 | string(&data["client"]), |
| 418 | string(&data["session"]), |
| 419 | expires, |
| 420 | family, |
| 421 | scope, |
| 422 | data["auth_time"].as_i64().unwrap_or(issued) |
| 423 | ], |
| 424 | )?; |
| 425 | if kind == "refresh" { |
| 426 | output["refresh_token"] = json!(token); |
| 427 | } |
| 428 | } |
| 429 | Ok(output) |
| 430 | } |
| 431 | |
| 432 | async fn handle(app: &App, request: Request) -> Result<Response> { |
| 433 | let auth = &app.auth; |
| 434 | let path = request.uri().path().to_owned(); |
| 435 | let method = request.method().clone(); |
| 436 | let headers = request.headers().clone(); |
| 437 | let query = fields(request.uri().query().unwrap_or_default())?; |
| 438 | let issuer = auth.origin.origin().ascii_serialization(); |
| 439 | if path == "/.well-known/openid-configuration" && method == Method::GET { |
| 440 | return Ok(axum::Json(json!({"issuer":issuer,"authorization_endpoint":format!("{issuer}/auth/oidc/authorize"),"token_endpoint":format!("{issuer}/auth/oidc/token"),"userinfo_endpoint":format!("{issuer}/auth/oidc/userinfo"),"jwks_uri":format!("{issuer}/auth/oidc/jwks"),"revocation_endpoint":format!("{issuer}/auth/oidc/revoke"),"response_types_supported":["code"],"response_modes_supported":["query"],"grant_types_supported":["authorization_code","refresh_token"],"subject_types_supported":["public"],"id_token_signing_alg_values_supported":["RS256"],"token_endpoint_auth_methods_supported":["client_secret_basic","client_secret_post"],"code_challenge_methods_supported":["S256"],"scopes_supported":SCOPES,"claims_supported":["sub","preferred_username","name","given_name","family_name","email","email_verified","groups","auth_time","nonce"]})).into_response()); |
| 441 | } |
| 442 | if path == "/auth/oidc/jwks" && method == Method::GET { |
| 443 | let (key, kid) = public_key(&auth.db.lock().unwrap())?; |
| 444 | let rsa = key.rsa()?; |
| 445 | return Ok(axum::Json(json!({"keys":[{"kty":"RSA","use":"sig","alg":"RS256","kid":kid,"n":URL_SAFE_NO_PAD.encode(rsa.n().to_vec()),"e":URL_SAFE_NO_PAD.encode(rsa.e().to_vec())}]})).into_response()); |
| 446 | } |
| 447 | if path == "/auth/oidc/authorize" && method == Method::GET { |
| 448 | let id = query |
| 449 | .get("client_id") |
| 450 | .map(String::as_str) |
| 451 | .unwrap_or_default(); |
| 452 | let redirect = query |
| 453 | .get("redirect_uri") |
| 454 | .map(String::as_str) |
| 455 | .unwrap_or_default(); |
| 456 | let config = client(&auth.db.lock().unwrap(), id)?; |
| 457 | if !array(&config["redirectUris"]) |
| 458 | .iter() |
| 459 | .any(|v| v.as_str() == Some(redirect)) |
| 460 | { |
| 461 | return Err(invalid("invalid_redirect_uri")); |
| 462 | } |
| 463 | if query.get("response_type").map(String::as_str) != Some("code") |
| 464 | || query.get("response_mode").is_some_and(|m| m != "query") |
| 465 | { |
| 466 | return Err(invalid("unsupported_response_type")); |
| 467 | } |
| 468 | let scope = query.get("scope").map(String::as_str).unwrap_or_default(); |
| 469 | if !scope.split_whitespace().any(|s| s == "openid") |
| 470 | || scope.split_whitespace().any(|s| !SCOPES.contains(&s)) |
| 471 | { |
| 472 | return Err(invalid("invalid_scope")); |
| 473 | } |
| 474 | let challenge = query.get("code_challenge"); |
| 475 | if challenge.is_some_and(|c| { |
| 476 | c.len() != 43 |
| 477 | || !c |
| 478 | .chars() |
| 479 | .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-') |
| 480 | }) || (challenge.is_some() |
| 481 | && query.get("code_challenge_method").map(String::as_str) != Some("S256")) |
| 482 | || (challenge.is_none() && query.contains_key("code_challenge_method")) |
| 483 | { |
| 484 | return Err(invalid("invalid_request")); |
| 485 | } |
| 486 | if query.values().any(|v| v.len() > 4096) { |
| 487 | return Err(invalid("invalid_request")); |
| 488 | } |
| 489 | let prompt = query.get("prompt").map(String::as_str).unwrap_or_default(); |
| 490 | if !matches!(prompt, "" | "none" | "login") { |
| 491 | return Err(invalid("invalid_request")); |
| 492 | } |
| 493 | let user = auth.session(&headers, "dashboard")?; |
| 494 | if guest::is_guest(&user) && !guests_allowed(id, &config) { |
| 495 | return Err(Error::new(403, "access_denied")); |
| 496 | } |
| 497 | let session = mcp::hash(&auth::cookie(&headers, "__Host-snow-session").unwrap_or_default()); |
| 498 | let auth_time: i64 = auth |
| 499 | .db |
| 500 | .lock() |
| 501 | .unwrap() |
| 502 | .query_row( |
| 503 | "SELECT auth_time FROM sessions WHERE hash=?", |
| 504 | [&session], |
| 505 | |r| r.get(0), |
| 506 | ) |
| 507 | .optional()? |
| 508 | .unwrap_or(0); |
| 509 | let max_age = query |
| 510 | .get("max_age") |
| 511 | .map(|n| { |
| 512 | n.parse::<i64>() |
| 513 | .ok() |
| 514 | .filter(|n| *n >= 0) |
| 515 | .ok_or_else(|| invalid("invalid_request")) |
| 516 | }) |
| 517 | .transpose()?; |
| 518 | let reauth_token = query |
| 519 | .get("snow_reauth") |
| 520 | .map(String::as_str) |
| 521 | .unwrap_or_default(); |
| 522 | let reauth_state = |
| 523 | auth::pending(&auth.db.lock().unwrap(), reauth_token, "oidc-reauth", false)?; |
| 524 | let reauthed = reauth_state["client"] == id |
| 525 | && reauth_state["redirect"] == redirect |
| 526 | && reauth_state["session"] |
| 527 | .as_str() |
| 528 | .is_some_and(|old| old != session); |
| 529 | let reauth = prompt == "login" && !reauthed |
| 530 | || max_age.is_some_and(|age| now() as i64 - auth_time > age); |
| 531 | let mut target = url::Url::parse(redirect)?; |
| 532 | if user.is_null() || reauth { |
| 533 | if prompt == "none" { |
| 534 | target |
| 535 | .query_pairs_mut() |
| 536 | .append_pair("error", "login_required"); |
| 537 | if let Some(state) = query.get("state") { |
| 538 | target.query_pairs_mut().append_pair("state", state); |
| 539 | } |
| 540 | return Ok((StatusCode::FOUND, [("location", target.to_string())]).into_response()); |
| 541 | } |
| 542 | let mut next = auth.origin.join(&request.uri().to_string())?; |
| 543 | if prompt == "login" { |
| 544 | let token = if reauth_state["client"] == id && reauth_state["redirect"] == redirect |
| 545 | { |
| 546 | reauth_token.to_owned() |
| 547 | } else { |
| 548 | auth::issue( |
| 549 | &auth.db.lock().unwrap(), |
| 550 | "oidc-reauth", |
| 551 | json!({"client":id,"redirect":redirect,"session":session}), |
| 552 | 900, |
| 553 | )? |
| 554 | }; |
| 555 | let pairs: Vec<_> = next |
| 556 | .query_pairs() |
| 557 | .filter(|(key, _)| key != "snow_reauth") |
| 558 | .map(|(k, v)| (k.into_owned(), v.into_owned())) |
| 559 | .collect(); |
| 560 | next.set_query(None); |
| 561 | next.query_pairs_mut() |
| 562 | .extend_pairs(pairs) |
| 563 | .append_pair("snow_reauth", &token); |
| 564 | } |
| 565 | return Ok(( |
| 566 | StatusCode::FOUND, |
| 567 | [( |
| 568 | "location", |
| 569 | format!( |
| 570 | "/sign-in?next={}", |
| 571 | encoded(&format!( |
| 572 | "{}?{}", |
| 573 | next.path(), |
| 574 | next.query().unwrap_or_default() |
| 575 | )) |
| 576 | ), |
| 577 | )], |
| 578 | ) |
| 579 | .into_response()); |
| 580 | } |
| 581 | if !array(&user["requiredActions"]).is_empty() { |
| 582 | return Ok((StatusCode::FOUND, [("location", "/account")]).into_response()); |
| 583 | } |
| 584 | if reauthed { |
| 585 | auth::pending(&auth.db.lock().unwrap(), reauth_token, "oidc-reauth", true)?; |
| 586 | } |
| 587 | let data = json!({"user":user["id"],"client":id,"redirect":redirect,"scope":scope,"challenge":challenge,"nonce":query.get("nonce"),"session":session,"auth_time":auth_time,"family":mcp::secret()}); |
| 588 | let code = auth::issue(&auth.db.lock().unwrap(), "oidc-code", data, 60)?; |
| 589 | target.query_pairs_mut().append_pair("code", &code); |
| 590 | if let Some(state) = query.get("state") { |
| 591 | target.query_pairs_mut().append_pair("state", state); |
| 592 | } |
| 593 | return Ok((StatusCode::FOUND, [("location", target.to_string())]).into_response()); |
| 594 | } |
| 595 | if path == "/auth/oidc/userinfo" && matches!(method, Method::GET | Method::POST) { |
| 596 | let token = headers |
| 597 | .get("authorization") |
| 598 | .and_then(|h| h.to_str().ok()) |
| 599 | .and_then(|h| h.strip_prefix("Bearer ")) |
| 600 | .ok_or_else(|| Error::new(401, "invalid_token"))?; |
| 601 | let db = auth.db.lock().unwrap(); |
| 602 | let data: Option<(String,String,String,String)> = db.query_row("SELECT user_id,client_id,session_hash,scope FROM oidc_tokens WHERE hash=? AND kind='access' AND expires>?",sql![mcp::hash(token),now() as i64],|r|Ok((r.get(0)?,r.get(1)?,r.get(2)?,r.get(3)?))).optional()?; |
| 603 | let (user, client_id, session, scope) = |
| 604 | data.ok_or_else(|| Error::new(401, "invalid_token"))?; |
| 605 | let profile = |
| 606 | eligible(&db, &user, &session).map_err(|_| Error::new(401, "invalid_token"))?; |
| 607 | let config = client(&db, &client_id)?; |
| 608 | if guest::is_guest(&profile) && !guests_allowed(&client_id, &config) { |
| 609 | return Err(Error::new(401, "invalid_token")); |
| 610 | } |
| 611 | return Ok(axum::Json(claims(&profile, &config, &scope)).into_response()); |
| 612 | } |
| 613 | if method != Method::POST || !matches!(path.as_str(), "/auth/oidc/token" | "/auth/oidc/revoke") |
| 614 | { |
| 615 | return Err(Error::new(404, "not_found")); |
| 616 | } |
| 617 | if headers |
| 618 | .get("content-type") |
| 619 | .and_then(|h| h.to_str().ok()) |
| 620 | .is_none_or(|t| t.split(';').next() != Some("application/x-www-form-urlencoded")) |
| 621 | { |
| 622 | return Err(invalid("invalid_request")); |
| 623 | } |
| 624 | let body = axum::body::to_bytes(request.into_body(), 16384).await?; |
| 625 | let form = fields(std::str::from_utf8(&body).map_err(|_| invalid("invalid_request"))?)?; |
| 626 | let mut db = auth.db.lock().unwrap(); |
| 627 | let tx = db.transaction()?; |
| 628 | let id = authenticated_client(&tx, &headers, &form)?; |
| 629 | if path.ends_with("/revoke") { |
| 630 | tx.execute("DELETE FROM oidc_tokens WHERE client_id=? AND family=(SELECT family FROM oidc_tokens WHERE hash=? AND client_id=?)",sql![id,mcp::hash(form.get("token").map(String::as_str).unwrap_or_default()),id])?; |
| 631 | tx.commit()?; |
| 632 | return Ok(StatusCode::OK.into_response()); |
| 633 | } |
| 634 | let grant = form |
| 635 | .get("grant_type") |
| 636 | .map(String::as_str) |
| 637 | .unwrap_or_default(); |
| 638 | let data = if grant == "authorization_code" { |
| 639 | let code = form.get("code").map(String::as_str).unwrap_or_default(); |
| 640 | let data = auth::pending(&tx, code, "oidc-code", false)?; |
| 641 | if data.is_null() |
| 642 | || data["client"] != id |
| 643 | || data["redirect"].as_str() != form.get("redirect_uri").map(String::as_str) |
| 644 | { |
| 645 | return Err(invalid("invalid_grant")); |
| 646 | } |
| 647 | if let Some(challenge) = data["challenge"].as_str() { |
| 648 | let verifier = form |
| 649 | .get("code_verifier") |
| 650 | .map(String::as_str) |
| 651 | .unwrap_or_default(); |
| 652 | if !(43..=128).contains(&verifier.len()) |
| 653 | || !verifier |
| 654 | .chars() |
| 655 | .all(|c| c.is_ascii_alphanumeric() || "-._~".contains(c)) |
| 656 | || URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes())) != challenge |
| 657 | { |
| 658 | return Err(invalid("invalid_grant")); |
| 659 | } |
| 660 | } |
| 661 | auth::pending(&tx, code, "oidc-code", true)?; |
| 662 | data |
| 663 | } else if grant == "refresh_token" { |
| 664 | let hash = mcp::hash( |
| 665 | form.get("refresh_token") |
| 666 | .map(String::as_str) |
| 667 | .unwrap_or_default(), |
| 668 | ); |
| 669 | let data: Option<(String,String,String,String,String,i64)> = tx.query_row("SELECT kind,user_id,session_hash,family,scope,auth_time FROM oidc_tokens WHERE hash=? AND client_id=? AND expires>?",sql![hash,id,now() as i64],|r|Ok((r.get(0)?,r.get(1)?,r.get(2)?,r.get(3)?,r.get(4)?,r.get(5)?))).optional()?; |
| 670 | let (kind, user, session, family, scope, auth_time) = |
| 671 | data.ok_or_else(|| invalid("invalid_grant"))?; |
| 672 | if kind == "refresh-used" { |
| 673 | tx.execute("DELETE FROM oidc_tokens WHERE family=?", [family])?; |
| 674 | tx.commit()?; |
| 675 | return Err(invalid("invalid_grant")); |
| 676 | } |
| 677 | if kind != "refresh" || form.get("scope").is_some_and(|s| s != &scope) { |
| 678 | return Err(invalid("invalid_grant")); |
| 679 | } |
| 680 | tx.execute( |
| 681 | "UPDATE oidc_tokens SET kind='refresh-used' WHERE hash=?", |
| 682 | [hash], |
| 683 | )?; |
| 684 | tx.execute( |
| 685 | "DELETE FROM oidc_tokens WHERE family=? AND kind='access'", |
| 686 | [&family], |
| 687 | )?; |
| 688 | json!({"user":user,"client":id,"session":session,"family":family,"scope":scope,"auth_time":auth_time}) |
| 689 | } else { |
| 690 | return Err(invalid("unsupported_grant_type")); |
| 691 | }; |
| 692 | tx.execute("DELETE FROM oidc_tokens WHERE expires<=?", [now() as i64])?; |
| 693 | let output = tokens(&tx, auth, &data, data["nonce"].as_str())?; |
| 694 | tx.commit()?; |
| 695 | Ok(axum::Json(output).into_response()) |
| 696 | } |
| 697 | |
| 698 | pub async fn route(State(app): State<Arc<App>>, request: Request) -> Response { |
| 699 | let mut response = match handle(&app, request).await { |
| 700 | Ok(response) => response, |
| 701 | Err(error) => { |
| 702 | if error.status >= 500 { |
| 703 | eprintln!("OIDC: {}", error.message); |
| 704 | } |
| 705 | ( |
| 706 | StatusCode::from_u16(error.status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), |
| 707 | axum::Json( |
| 708 | json!({"error":if error.status>=500 {"server_error"} else {&error.message}}), |
| 709 | ), |
| 710 | ) |
| 711 | .into_response() |
| 712 | } |
| 713 | }; |
| 714 | response |
| 715 | .headers_mut() |
| 716 | .insert("cache-control", "no-store".parse().unwrap()); |
| 717 | response |
| 718 | .headers_mut() |
| 719 | .insert("pragma", "no-cache".parse().unwrap()); |
| 720 | response |
| 721 | } |