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