| 1 | use crate::*; |
| 2 | use base64::{ |
| 3 | Engine, |
| 4 | engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}, |
| 5 | }; |
| 6 | use rusqlite::{Connection, OpenFlags, OptionalExtension}; |
| 7 | use sha2::{Digest, Sha256}; |
| 8 | use std::os::unix::fs::{DirBuilderExt, PermissionsExt}; |
| 9 | |
| 10 | const CATALOGS: &[(&str, &str, &[&str])] = &[ |
| 11 | ( |
| 12 | "observability", |
| 13 | "Logs and traces", |
| 14 | &["observability:read", "offline_access"], |
| 15 | ), |
| 16 | ( |
| 17 | "agents", |
| 18 | "Local agents", |
| 19 | &["sessions:read", "sessions:write", "offline_access"], |
| 20 | ), |
| 21 | ( |
| 22 | "shale", |
| 23 | "Shale", |
| 24 | &["shale:read", "shale:write", "offline_access"], |
| 25 | ), |
| 26 | ]; |
| 27 | |
| 28 | pub struct Store { |
| 29 | pub(crate) db: Mutex<Connection>, |
| 30 | pub origin: url::Url, |
| 31 | } |
| 32 | pub(crate) fn secret() -> String { |
| 33 | URL_SAFE_NO_PAD.encode(rand::random::<[u8; 32]>()) |
| 34 | } |
| 35 | pub(crate) fn hash(value: &str) -> String { |
| 36 | format!("{:x}", Sha256::digest(value.as_bytes())) |
| 37 | } |
| 38 | pub(crate) fn get(db: &Connection, key: &str) -> Result<Value> { |
| 39 | let row: Option<String> = db |
| 40 | .query_row( |
| 41 | "SELECT value FROM records WHERE key=? AND (expires=0 OR expires>?)", |
| 42 | rusqlite::params![key, now() as i64], |
| 43 | |r| r.get(0), |
| 44 | ) |
| 45 | .optional()?; |
| 46 | Ok(row |
| 47 | .map(|s| serde_json::from_str(&s)) |
| 48 | .transpose()? |
| 49 | .unwrap_or(Value::Null)) |
| 50 | } |
| 51 | pub(crate) fn put(db: &Connection, key: &str, value: &Value, ttl: i64) -> Result<()> { |
| 52 | db.execute( |
| 53 | "DELETE FROM records WHERE expires>0 AND expires<=?", |
| 54 | [now() as i64], |
| 55 | )?; |
| 56 | let count: i64 = db.query_row("SELECT count(*) FROM records", [], |r| r.get(0))?; |
| 57 | if count >= 16384 && get(db, key)?.is_null() { |
| 58 | return Err(Error::new( |
| 59 | 503, |
| 60 | "The connection store is full. Remove an unused connection and retry.", |
| 61 | )); |
| 62 | } |
| 63 | db.execute( |
| 64 | "INSERT OR REPLACE INTO records VALUES (?,?,?)", |
| 65 | rusqlite::params![ |
| 66 | key, |
| 67 | value.to_string(), |
| 68 | if ttl == 0 { 0 } else { now() as i64 + ttl } |
| 69 | ], |
| 70 | )?; |
| 71 | Ok(()) |
| 72 | } |
| 73 | pub(crate) fn delete(db: &Connection, key: &str) -> Result<()> { |
| 74 | db.execute("DELETE FROM records WHERE key=?", [key])?; |
| 75 | Ok(()) |
| 76 | } |
| 77 | pub(crate) fn revoke(db: &Connection, grant: &str) -> Result<()> { |
| 78 | db.execute( |
| 79 | "DELETE FROM records WHERE key=? OR json_extract(value,'$.grant')=?", |
| 80 | rusqlite::params![format!("grant:{grant}"), grant], |
| 81 | )?; |
| 82 | Ok(()) |
| 83 | } |
| 84 | pub(crate) fn list(db: &Connection, prefix: &str) -> Result<Vec<Value>> { |
| 85 | let mut query = db.prepare( |
| 86 | "SELECT value FROM records WHERE substr(key,1,?)=? AND (expires=0 OR expires>?)", |
| 87 | )?; |
| 88 | let rows = query.query_map(rusqlite::params![prefix.len(), prefix, now() as i64], |r| { |
| 89 | r.get::<_, String>(0) |
| 90 | })?; |
| 91 | rows.map(|r| Ok(serde_json::from_str(&r?)?)).collect() |
| 92 | } |
| 93 | fn fail(code: &str) -> Error { |
| 94 | Error::new(400, code) |
| 95 | } |
| 96 | fn redirect(value: &str) -> Result<url::Url> { |
| 97 | let url = url::Url::parse(value).map_err(|_| fail("invalid_redirect_uri"))?; |
| 98 | let loopback = matches!(url.host_str(), Some("localhost" | "127.0.0.1" | "[::1]")); |
| 99 | if value.len() > 4096 |
| 100 | || !url.username().is_empty() |
| 101 | || url.password().is_some() |
| 102 | || url.fragment().is_some() |
| 103 | || url.host_str().is_none() |
| 104 | || !(url.scheme() == "https" || url.scheme() == "http" && loopback) |
| 105 | { |
| 106 | return Err(fail("invalid_redirect_uri")); |
| 107 | } |
| 108 | Ok(url) |
| 109 | } |
| 110 | impl Store { |
| 111 | pub fn new(data: &std::path::Path, origin: &str) -> Result<Self> { |
| 112 | let origin = url::Url::parse(origin)?; |
| 113 | if origin.path() != "/" |
| 114 | || origin.query().is_some() |
| 115 | || origin.fragment().is_some() |
| 116 | || !origin.username().is_empty() |
| 117 | || origin.password().is_some() |
| 118 | || origin.scheme() != "https" |
| 119 | && !(origin.scheme() == "http" |
| 120 | && matches!(origin.host_str(), Some("localhost" | "127.0.0.1" | "[::1]"))) |
| 121 | { |
| 122 | return Err(fail("Set an HTTPS origin for MCP connections.")); |
| 123 | } |
| 124 | std::fs::DirBuilder::new() |
| 125 | .recursive(true) |
| 126 | .mode(0o700) |
| 127 | .create(data)?; |
| 128 | let path = data.join("connections.sqlite"); |
| 129 | let db = Connection::open_with_flags( |
| 130 | &path, |
| 131 | OpenFlags::SQLITE_OPEN_READ_WRITE |
| 132 | | OpenFlags::SQLITE_OPEN_CREATE |
| 133 | | OpenFlags::SQLITE_OPEN_NOFOLLOW, |
| 134 | )?; |
| 135 | std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600))?; |
| 136 | db.execute_batch("PRAGMA journal_mode=WAL; PRAGMA busy_timeout=5000; CREATE TABLE IF NOT EXISTS records (key TEXT PRIMARY KEY, value TEXT NOT NULL, expires INTEGER NOT NULL);")?; |
| 137 | Ok(Self { |
| 138 | db: Mutex::new(db), |
| 139 | origin, |
| 140 | }) |
| 141 | } |
| 142 | pub fn resource(&self, catalog: &str) -> String { |
| 143 | self.origin |
| 144 | .join(&format!("mcp/{catalog}")) |
| 145 | .unwrap() |
| 146 | .to_string() |
| 147 | } |
| 148 | fn client( |
| 149 | &self, |
| 150 | db: &Connection, |
| 151 | input: &HashMap<String, String>, |
| 152 | headers: &HeaderMap, |
| 153 | ) -> Result<Value> { |
| 154 | let basic = headers.get("authorization").and_then(|v| v.to_str().ok()); |
| 155 | let (id, credential) = if let Some(basic) = basic { |
| 156 | if input.contains_key("client_secret") { |
| 157 | return Err(fail("invalid_client")); |
| 158 | } |
| 159 | let bytes = STANDARD |
| 160 | .decode( |
| 161 | basic |
| 162 | .strip_prefix("Basic ") |
| 163 | .ok_or_else(|| fail("invalid_client"))?, |
| 164 | ) |
| 165 | .map_err(|_| fail("invalid_client"))?; |
| 166 | let value = String::from_utf8(bytes).map_err(|_| fail("invalid_client"))?; |
| 167 | let (id, credential) = value |
| 168 | .split_once(':') |
| 169 | .ok_or_else(|| fail("invalid_client"))?; |
| 170 | (id.to_owned(), Some(credential.to_owned())) |
| 171 | } else { |
| 172 | ( |
| 173 | input.get("client_id").cloned().unwrap_or_default(), |
| 174 | input.get("client_secret").cloned(), |
| 175 | ) |
| 176 | }; |
| 177 | if input.get("client_id").is_some_and(|v| v != &id) { |
| 178 | return Err(fail("invalid_client")); |
| 179 | } |
| 180 | let client = get(db, &format!("client:{id}"))?; |
| 181 | let method = string(&client["token_endpoint_auth_method"]); |
| 182 | let valid = match method { |
| 183 | "none" => basic.is_none() && credential.is_none(), |
| 184 | "client_secret_basic" => { |
| 185 | basic.is_some() |
| 186 | && credential.as_ref().is_some_and(|v| { |
| 187 | hash(v) |
| 188 | .as_bytes() |
| 189 | .ct_eq(string(&client["secret_hash"]).as_bytes()) |
| 190 | .unwrap_u8() |
| 191 | == 1 |
| 192 | }) |
| 193 | } |
| 194 | "client_secret_post" => { |
| 195 | basic.is_none() |
| 196 | && credential.as_ref().is_some_and(|v| { |
| 197 | hash(v) |
| 198 | .as_bytes() |
| 199 | .ct_eq(string(&client["secret_hash"]).as_bytes()) |
| 200 | .unwrap_u8() |
| 201 | == 1 |
| 202 | }) |
| 203 | } |
| 204 | _ => false, |
| 205 | }; |
| 206 | if valid { |
| 207 | Ok(client) |
| 208 | } else { |
| 209 | Err(fail("invalid_client")) |
| 210 | } |
| 211 | } |
| 212 | fn issue(&self, db: &Connection, grant: &Value) -> Result<Value> { |
| 213 | let access = secret(); |
| 214 | put( |
| 215 | db, |
| 216 | &format!("access:{}", hash(&access)), |
| 217 | &json!({"grant":grant["id"],"resource":grant["resource"]}), |
| 218 | 3600, |
| 219 | )?; |
| 220 | let mut response = json!({"access_token":access,"token_type":"Bearer","expires_in":3600,"scope":array(&grant["scopes"]).iter().map(string).collect::<Vec<_>>().join(" ")}); |
| 221 | if array(&grant["scopes"]) |
| 222 | .iter() |
| 223 | .any(|s| s == "offline_access") |
| 224 | { |
| 225 | let refresh = secret(); |
| 226 | put( |
| 227 | db, |
| 228 | &format!("refresh:{}", hash(&refresh)), |
| 229 | &json!({"grant":grant["id"]}), |
| 230 | 30 * 86400, |
| 231 | )?; |
| 232 | response["refresh_token"] = json!(refresh); |
| 233 | } |
| 234 | Ok(response) |
| 235 | } |
| 236 | pub fn authenticate(&self, headers: &HeaderMap, resource: &str) -> Result<Value> { |
| 237 | let token = headers |
| 238 | .get("authorization") |
| 239 | .and_then(|v| v.to_str().ok()) |
| 240 | .and_then(|s| s.strip_prefix("Bearer ")) |
| 241 | .filter(|s| !s.is_empty() && s.len() <= 256) |
| 242 | .ok_or_else(|| Error::new(401, "invalid_token"))?; |
| 243 | let db = self.db.lock().unwrap(); |
| 244 | let access = get(&db, &format!("access:{}", hash(token)))?; |
| 245 | let grant = get(&db, &format!("grant:{}", string(&access["grant"])))?; |
| 246 | if grant.is_null() || access["resource"] != resource || grant["resource"] != resource { |
| 247 | return Err(Error::new(401, "invalid_token")); |
| 248 | } |
| 249 | Ok(grant) |
| 250 | } |
| 251 | fn exchange( |
| 252 | &self, |
| 253 | db: &Connection, |
| 254 | input: &HashMap<String, String>, |
| 255 | headers: &HeaderMap, |
| 256 | ) -> Result<(String, Value)> { |
| 257 | let client = self.client(db, input, headers)?; |
| 258 | let (key, grant) = match input.get("grant_type").map(String::as_str) { |
| 259 | Some("authorization_code") => { |
| 260 | let key = format!( |
| 261 | "code:{}", |
| 262 | hash(input.get("code").map(String::as_str).unwrap_or_default()) |
| 263 | ); |
| 264 | let code = get(db, &key)?; |
| 265 | let verifier = input |
| 266 | .get("code_verifier") |
| 267 | .map(String::as_str) |
| 268 | .unwrap_or_default(); |
| 269 | if code.is_null() |
| 270 | || code["client"] != client["client_id"] |
| 271 | || input.get("redirect_uri").map(String::as_str) != code["redirect"].as_str() |
| 272 | || !(43..=128).contains(&verifier.len()) |
| 273 | || !verifier |
| 274 | .bytes() |
| 275 | .all(|b| b.is_ascii_alphanumeric() || b"-._~".contains(&b)) |
| 276 | || URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes())) |
| 277 | != string(&code["challenge"]) |
| 278 | { |
| 279 | return Err(fail("invalid_grant")); |
| 280 | } |
| 281 | if input |
| 282 | .get("resource") |
| 283 | .is_some_and(|r| code["resource"] != r.as_str()) |
| 284 | { |
| 285 | return Err(fail("invalid_target")); |
| 286 | } |
| 287 | let grant = get(db, &format!("grant:{}", string(&code["grant"])))?; |
| 288 | if grant.is_null() { |
| 289 | return Err(fail("invalid_grant")); |
| 290 | } |
| 291 | (key, grant) |
| 292 | } |
| 293 | Some("refresh_token") => { |
| 294 | let fingerprint = hash( |
| 295 | input |
| 296 | .get("refresh_token") |
| 297 | .map(String::as_str) |
| 298 | .unwrap_or_default(), |
| 299 | ); |
| 300 | let key = format!("refresh:{fingerprint}"); |
| 301 | let token = get(db, &key)?; |
| 302 | let used = get(db, &format!("used:{fingerprint}"))?; |
| 303 | if !used.is_null() { |
| 304 | let grant = get(db, &format!("grant:{}", string(&used["grant"])))?; |
| 305 | if grant["client"] == client["client_id"] { |
| 306 | revoke(db, string(&used["grant"]))?; |
| 307 | } |
| 308 | return Err(fail("invalid_grant")); |
| 309 | } |
| 310 | let grant = get(db, &format!("grant:{}", string(&token["grant"])))?; |
| 311 | if grant.is_null() || grant["client"] != client["client_id"] { |
| 312 | return Err(fail("invalid_grant")); |
| 313 | } |
| 314 | if input |
| 315 | .get("resource") |
| 316 | .is_some_and(|r| grant["resource"] != r.as_str()) |
| 317 | { |
| 318 | return Err(fail("invalid_target")); |
| 319 | } |
| 320 | if input.get("scope").is_some_and(|s| { |
| 321 | s.split_whitespace() |
| 322 | .collect::<std::collections::HashSet<_>>() |
| 323 | != array(&grant["scopes"]).iter().map(string).collect() |
| 324 | }) { |
| 325 | return Err(fail("invalid_scope")); |
| 326 | } |
| 327 | (key, grant) |
| 328 | } |
| 329 | _ => return Err(fail("unsupported_grant_type")), |
| 330 | }; |
| 331 | Ok((key, grant)) |
| 332 | } |
| 333 | fn oauth( |
| 334 | &self, |
| 335 | path: &str, |
| 336 | method: &Method, |
| 337 | input: &HashMap<String, String>, |
| 338 | headers: &HeaderMap, |
| 339 | body: Value, |
| 340 | ) -> Result<Response> { |
| 341 | let mut db = self.db.lock().unwrap(); |
| 342 | let tx = db.transaction()?; |
| 343 | let value = match (path, method) { |
| 344 | ("register", &Method::POST) => { |
| 345 | let name = body["client_name"] |
| 346 | .as_str() |
| 347 | .filter(|s| !s.is_empty() && s.len() <= 128) |
| 348 | .ok_or_else(|| fail("invalid_client_metadata"))?; |
| 349 | let uris = body["redirect_uris"] |
| 350 | .as_array() |
| 351 | .filter(|a| !a.is_empty() && a.len() <= 8) |
| 352 | .ok_or_else(|| fail("invalid_client_metadata"))?; |
| 353 | for uri in uris { |
| 354 | redirect(uri.as_str().ok_or_else(|| fail("invalid_redirect_uri"))?)?; |
| 355 | } |
| 356 | let auth = body["token_endpoint_auth_method"] |
| 357 | .as_str() |
| 358 | .unwrap_or("none"); |
| 359 | if !["none", "client_secret_post", "client_secret_basic"].contains(&auth) |
| 360 | || body.get("grant_types").is_some_and(|v| { |
| 361 | v.as_array().is_none_or(|types| { |
| 362 | types.is_empty() |
| 363 | || types |
| 364 | .iter() |
| 365 | .any(|s| s != "authorization_code" && s != "refresh_token") |
| 366 | }) |
| 367 | }) |
| 368 | || body |
| 369 | .get("response_types") |
| 370 | .is_some_and(|v| v != &json!(["code"])) |
| 371 | { |
| 372 | return Err(fail("invalid_client_metadata")); |
| 373 | } |
| 374 | if list(&tx, "client:")?.len() >= 4096 { |
| 375 | return Err(Error::new(429, "too_many_clients")); |
| 376 | } |
| 377 | let id = uuid::Uuid::new_v4().to_string(); |
| 378 | let credential = secret(); |
| 379 | let mut client = json!({"client_id":id,"client_name":name,"redirect_uris":uris,"token_endpoint_auth_method":auth,"grant_types":["authorization_code","refresh_token"],"response_types":["code"],"client_id_issued_at":now() as i64}); |
| 380 | if auth != "none" { |
| 381 | client["secret_hash"] = json!(hash(&credential)); |
| 382 | } |
| 383 | put(&tx, &format!("client:{id}"), &client, 0)?; |
| 384 | client.as_object_mut().unwrap().remove("secret_hash"); |
| 385 | if auth != "none" { |
| 386 | client["client_secret"] = json!(credential); |
| 387 | client["client_secret_expires_at"] = json!(0); |
| 388 | } |
| 389 | tx.commit()?; |
| 390 | return Ok((StatusCode::CREATED, axum::Json(client)).into_response()); |
| 391 | } |
| 392 | ("authorize", &Method::GET) => { |
| 393 | let id = input |
| 394 | .get("client_id") |
| 395 | .map(String::as_str) |
| 396 | .unwrap_or_default(); |
| 397 | let client = get(&tx, &format!("client:{id}"))?; |
| 398 | let uri = input |
| 399 | .get("redirect_uri") |
| 400 | .map(String::as_str) |
| 401 | .unwrap_or_default(); |
| 402 | let challenge = input |
| 403 | .get("code_challenge") |
| 404 | .map(String::as_str) |
| 405 | .unwrap_or_default(); |
| 406 | let resource = input |
| 407 | .get("resource") |
| 408 | .map(String::as_str) |
| 409 | .unwrap_or_default(); |
| 410 | let scopes_allowed = CATALOGS |
| 411 | .iter() |
| 412 | .find(|(id, _, _)| resource == self.resource(id)) |
| 413 | .map(|(_, _, scopes)| *scopes) |
| 414 | .ok_or_else(|| fail("invalid_target"))?; |
| 415 | let scopes: Vec<_> = input |
| 416 | .get("scope") |
| 417 | .map(String::as_str) |
| 418 | .unwrap_or(scopes_allowed[0]) |
| 419 | .split_whitespace() |
| 420 | .collect(); |
| 421 | if client.is_null() || !array(&client["redirect_uris"]).iter().any(|v| v == uri) { |
| 422 | return Err(fail("invalid_redirect_uri")); |
| 423 | } |
| 424 | if input.get("response_type").map(String::as_str) != Some("code") |
| 425 | || input.get("code_challenge_method").map(String::as_str) != Some("S256") |
| 426 | || challenge.len() != 43 |
| 427 | || !challenge |
| 428 | .bytes() |
| 429 | .all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'_') |
| 430 | { |
| 431 | return Err(fail("invalid_request")); |
| 432 | } |
| 433 | if !scopes.contains(&scopes_allowed[0]) |
| 434 | || scopes.iter().any(|s| !scopes_allowed.contains(s)) |
| 435 | { |
| 436 | return Err(fail("invalid_scope")); |
| 437 | } |
| 438 | let pending = secret(); |
| 439 | put( |
| 440 | &tx, |
| 441 | &format!("pending:{}", hash(&pending)), |
| 442 | &json!({"client":id,"redirect":uri,"challenge":challenge,"resource":resource,"scopes":scopes,"state":input.get("state"),"owner":null}), |
| 443 | 600, |
| 444 | )?; |
| 445 | tx.commit()?; |
| 446 | return Ok(( |
| 447 | StatusCode::FOUND, |
| 448 | [("location", format!("/connect/{}", encoded(&pending)))], |
| 449 | ) |
| 450 | .into_response()); |
| 451 | } |
| 452 | ("token", &Method::POST) => { |
| 453 | let (key, grant) = match self.exchange(&tx, input, headers) { |
| 454 | Ok(exchange) => exchange, |
| 455 | Err(error) => { |
| 456 | tx.commit()?; |
| 457 | return Err(error); |
| 458 | } |
| 459 | }; |
| 460 | delete(&tx, &key)?; |
| 461 | if let Some(fingerprint) = key.strip_prefix("refresh:") { |
| 462 | put( |
| 463 | &tx, |
| 464 | &format!("used:{fingerprint}"), |
| 465 | &json!({"grant":grant["id"]}), |
| 466 | 30 * 86400, |
| 467 | )?; |
| 468 | } |
| 469 | self.issue(&tx, &grant)? |
| 470 | } |
| 471 | ("revoke", &Method::POST) => { |
| 472 | let client = self.client(&tx, input, headers)?; |
| 473 | let fingerprint = hash(input.get("token").map(String::as_str).unwrap_or_default()); |
| 474 | for kind in ["access", "refresh", "used"] { |
| 475 | let token = get(&tx, &format!("{kind}:{fingerprint}"))?; |
| 476 | let grant = get(&tx, &format!("grant:{}", string(&token["grant"])))?; |
| 477 | if !grant.is_null() && grant["client"] == client["client_id"] { |
| 478 | revoke(&tx, string(&grant["id"]))?; |
| 479 | } |
| 480 | } |
| 481 | tx.commit()?; |
| 482 | return Ok(StatusCode::OK.into_response()); |
| 483 | } |
| 484 | _ => return Err(Error::new(404, "No endpoint here.")), |
| 485 | }; |
| 486 | tx.commit()?; |
| 487 | Ok(axum::Json(value).into_response()) |
| 488 | } |
| 489 | } |
| 490 | |
| 491 | #[derive(Clone)] |
| 492 | pub(crate) struct Grant(pub Value); |
| 493 | pub(crate) fn active_owner(app: &App, grant: &Value) -> Result<bool> { |
| 494 | let profile = match auth::user(&app.auth.db.lock().unwrap(), string(&grant["user"])) { |
| 495 | Ok(profile) => profile, |
| 496 | Err(error) if error.status == 404 => return Ok(false), |
| 497 | Err(error) => return Err(error), |
| 498 | }; |
| 499 | Ok(profile["enabled"] == true |
| 500 | && !guest::is_guest(&profile) |
| 501 | && (grant["resource"] != app.mcp.resource("observability") |
| 502 | || array(&profile["groups"]) |
| 503 | .iter() |
| 504 | .any(|role| role["name"] == "infra-admin"))) |
| 505 | } |
| 506 | pub fn router<H: rmcp::ServerHandler>( |
| 507 | app: Arc<App>, |
| 508 | catalog: &str, |
| 509 | handler: impl Fn() -> std::result::Result<H, std::io::Error> + Send + Sync + 'static, |
| 510 | ) -> Router { |
| 511 | use rmcp::transport::streamable_http_server::{ |
| 512 | StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager, |
| 513 | }; |
| 514 | let resource = app.mcp.resource(catalog); |
| 515 | let metadata = format!( |
| 516 | "{}.well-known/oauth-protected-resource/mcp/{catalog}", |
| 517 | app.mcp.origin.as_str() |
| 518 | ); |
| 519 | let mut config = StreamableHttpServerConfig::default(); |
| 520 | config.legacy_session_mode = false; |
| 521 | config.json_response = true; |
| 522 | config.allowed_hosts = |
| 523 | vec![app.mcp.origin[url::Position::BeforeHost..url::Position::AfterPort].to_owned()]; |
| 524 | config.allowed_origins = vec![format!( |
| 525 | "{}://{}:{}", |
| 526 | app.mcp.origin.scheme(), |
| 527 | app.mcp.origin.host_str().unwrap(), |
| 528 | app.mcp.origin.port_or_known_default().unwrap() |
| 529 | )]; |
| 530 | let service = |
| 531 | StreamableHttpService::new(handler, Arc::new(LocalSessionManager::default()), config); |
| 532 | Router::new() |
| 533 | .nest_service(&format!("/mcp/{catalog}"), service) |
| 534 | .layer(axum::middleware::from_fn( |
| 535 | move |mut request: Request, next: axum::middleware::Next| { |
| 536 | let (app, resource, metadata) = (app.clone(), resource.clone(), metadata.clone()); |
| 537 | async move { |
| 538 | if request.headers().get("origin").is_some_and(|origin| { |
| 539 | origin.to_str().ok() |
| 540 | != Some(app.mcp.origin.origin().ascii_serialization().as_str()) |
| 541 | }) { |
| 542 | return Error::new(403, "This origin cannot use the connector.") |
| 543 | .into_response(); |
| 544 | } |
| 545 | match app |
| 546 | .mcp |
| 547 | .authenticate(request.headers(), &resource) |
| 548 | .and_then(|grant| { |
| 549 | if active_owner(&app, &grant)? { |
| 550 | Ok(grant) |
| 551 | } else { |
| 552 | Err(Error::new(401, "invalid_token")) |
| 553 | } |
| 554 | }) { |
| 555 | Ok(grant) => { |
| 556 | request.extensions_mut().insert(Grant(grant)); |
| 557 | if let Some(value) = request.headers_mut().get_mut("authorization") { |
| 558 | value.set_sensitive(true); |
| 559 | } |
| 560 | next.run(request).await |
| 561 | } |
| 562 | Err(_) => ( |
| 563 | StatusCode::UNAUTHORIZED, |
| 564 | [( |
| 565 | "www-authenticate", |
| 566 | format!("Bearer resource_metadata=\"{metadata}\""), |
| 567 | )], |
| 568 | "This connection expired. Connect again.", |
| 569 | ) |
| 570 | .into_response(), |
| 571 | } |
| 572 | } |
| 573 | }, |
| 574 | )) |
| 575 | } |
| 576 | |
| 577 | pub fn public(path: &str) -> bool { |
| 578 | path == "/.well-known/openid-configuration" |
| 579 | || path.starts_with("/oauth/") |
| 580 | || CATALOGS.iter().any(|(id, _, _)| { |
| 581 | path == format!("/mcp/{id}") || path.starts_with(&format!("/mcp/{id}/")) |
| 582 | }) |
| 583 | || path.starts_with("/.well-known/oauth-") |
| 584 | || path == "/pairing" |
| 585 | || path == "/agent/connect" |
| 586 | || matches!( |
| 587 | path, |
| 588 | "/agent/install.sh" | "/agent/install.ps1" | "/agent/setup.mjs" | "/agent/relay.mjs" |
| 589 | ) |
| 590 | || path.starts_with("/api/v1/") |
| 591 | } |
| 592 | pub async fn oauth(State(app): State<Arc<App>>, request: Request) -> Response { |
| 593 | if request.uri().path().starts_with("/oauth/shale/") { |
| 594 | return shale::oauth(app, request).await; |
| 595 | } |
| 596 | let result = async { |
| 597 | let path = request.uri().path().to_owned(); |
| 598 | if path=="/.well-known/oauth-authorization-server" && request.method()==Method::GET { |
| 599 | let base = app.mcp.origin.as_str(); |
| 600 | let scopes: std::collections::BTreeSet<_> = CATALOGS.iter().flat_map(|(_, _, scopes)| scopes.iter()).collect(); |
| 601 | return Ok(axum::Json(json!({"issuer":base,"authorization_endpoint":format!("{base}oauth/authorize"),"token_endpoint":format!("{base}oauth/token"),"registration_endpoint":format!("{base}oauth/register"),"revocation_endpoint":format!("{base}oauth/revoke"),"response_types_supported":["code"],"grant_types_supported":["authorization_code","refresh_token"],"token_endpoint_auth_methods_supported":["none","client_secret_post","client_secret_basic"],"code_challenge_methods_supported":["S256"],"scopes_supported":scopes})).into_response()); |
| 602 | } |
| 603 | if request.method() == Method::GET { |
| 604 | for (catalog, _, scopes) in CATALOGS { |
| 605 | if path == format!("/.well-known/oauth-protected-resource/mcp/{catalog}") { |
| 606 | return Ok(axum::Json(json!({"resource":app.mcp.resource(catalog),"authorization_servers":[app.mcp.origin.as_str()],"scopes_supported":scopes,"bearer_methods_supported":["header"]})).into_response()); |
| 607 | } |
| 608 | } |
| 609 | } |
| 610 | let method = request.method().clone(); |
| 611 | let headers = request.headers().clone(); |
| 612 | let query = request.uri().query().unwrap_or_default().to_owned(); |
| 613 | let bytes = axum::body::to_bytes(request.into_body(),65536).await.map_err(|_| fail("invalid_request"))?; |
| 614 | let registration = path=="/oauth/register"; |
| 615 | let encoded = if method==Method::GET {query.as_bytes()} else {&bytes}; |
| 616 | if method==Method::POST && !headers.get("content-type").and_then(|v|v.to_str().ok()).is_some_and(|s| s.split(';').next()==Some(if registration {"application/json"} else {"application/x-www-form-urlencoded"})) { return Err(fail("invalid_request")); } |
| 617 | let mut input = HashMap::new(); |
| 618 | if !registration { for (key,value) in url::form_urlencoded::parse(encoded) { if input.insert(key.into_owned(),value.into_owned()).is_some() {return Err(fail("invalid_request"));} } } |
| 619 | let body = if registration {serde_json::from_slice(&bytes).map_err(|_| fail("invalid_client_metadata"))?} else {Value::Null}; |
| 620 | if path == "/oauth/token" && method == Method::POST { |
| 621 | let (_, grant) = app.mcp.exchange(&app.mcp.db.lock().unwrap(), &input, &headers)?; |
| 622 | if !active_owner(&app, &grant)? { |
| 623 | revoke(&app.mcp.db.lock().unwrap(), string(&grant["id"]))?; |
| 624 | return Err(fail("invalid_grant")); |
| 625 | } |
| 626 | } |
| 627 | app.mcp.oauth(path.trim_start_matches("/oauth/"),&method,&input,&headers,body) |
| 628 | }.await; |
| 629 | let mut response = match result { |
| 630 | Ok(response) => response, |
| 631 | Err(error) => ( |
| 632 | StatusCode::from_u16(if error.message == "invalid_client" { |
| 633 | 401 |
| 634 | } else { |
| 635 | error.status |
| 636 | }) |
| 637 | .unwrap_or(StatusCode::BAD_REQUEST), |
| 638 | axum::Json( |
| 639 | json!({"error":if error.status >= 500 {"server_error"} else {&error.message}}), |
| 640 | ), |
| 641 | ) |
| 642 | .into_response(), |
| 643 | }; |
| 644 | response |
| 645 | .headers_mut() |
| 646 | .insert("cache-control", "no-store".parse().unwrap()); |
| 647 | response |
| 648 | .headers_mut() |
| 649 | .insert("pragma", "no-cache".parse().unwrap()); |
| 650 | response |
| 651 | } |
| 652 | |
| 653 | fn chosen_resources(body: &Value, resources: &[Value], shale: bool) -> Result<Value> { |
| 654 | if shale && body["resources"] == "all" { |
| 655 | return Ok(json!("all")); |
| 656 | } |
| 657 | let chosen = body["resources"] |
| 658 | .as_array() |
| 659 | .filter(|items| !items.is_empty() && items.len() <= resources.len()) |
| 660 | .ok_or_else(|| Error::new(400, "Choose each available resource once."))?; |
| 661 | if chosen.iter().enumerate().any(|(index, item)| { |
| 662 | !resources.iter().any(|resource| item == &resource["id"]) || chosen[..index].contains(item) |
| 663 | }) { |
| 664 | return Err(Error::new( |
| 665 | 403, |
| 666 | "Choose resources available to your account.", |
| 667 | )); |
| 668 | } |
| 669 | Ok(json!(chosen)) |
| 670 | } |
| 671 | |
| 672 | pub async fn manage( |
| 673 | app: Arc<App>, |
| 674 | method: &Method, |
| 675 | parts: &[&str], |
| 676 | me: &Value, |
| 677 | body: Value, |
| 678 | headers: &HeaderMap, |
| 679 | ) -> Result<Response> { |
| 680 | if method != Method::GET |
| 681 | && (me["viewing"] == true |
| 682 | || headers.get("origin").and_then(|h| h.to_str().ok()) |
| 683 | != Some(app.mcp.origin.origin().ascii_serialization().as_str())) |
| 684 | { |
| 685 | return Err(Error::new( |
| 686 | 403, |
| 687 | "Open MCP settings from your own signed-in account.", |
| 688 | )); |
| 689 | } |
| 690 | let owner = users::self_user(&app, me).await?; |
| 691 | if owner["enabled"] != true { |
| 692 | return Err(Error::new(403, "This account is disabled.")); |
| 693 | } |
| 694 | let owner_id = string(&owner["id"]); |
| 695 | if parts == ["shale"] { |
| 696 | return shale::manage(app, method, &owner, &body).await; |
| 697 | } |
| 698 | if parts == ["relay", "live"] && method == Method::GET { |
| 699 | let owner = owner_id.to_owned(); |
| 700 | let stream = WatchStream::new(app.relay.changes.subscribe()).map(move |_| { |
| 701 | let machines = relay::machines(&app.mcp.db.lock().unwrap(), &owner) |
| 702 | .map(|machines| app.relay.view(machines, None)); |
| 703 | machines |
| 704 | .map(|machines| Event::default().json_data(machines).unwrap()) |
| 705 | .map_err(|error| std::io::Error::other(error.message)) |
| 706 | }); |
| 707 | return Ok(Sse::new(stream) |
| 708 | .keep_alive(axum::response::sse::KeepAlive::default()) |
| 709 | .into_response()); |
| 710 | } |
| 711 | let consent_resource = if let ["connections", id] = parts |
| 712 | && method != Method::DELETE |
| 713 | { |
| 714 | let grant = get(&app.mcp.db.lock().unwrap(), &format!("grant:{id}"))?; |
| 715 | if grant["user"] != owner_id { |
| 716 | return Err(Error::new(404, "No connection with that ID.")); |
| 717 | } |
| 718 | string(&grant["resource"]).to_owned() |
| 719 | } else if let ["consent", id] = parts { |
| 720 | let pending = get( |
| 721 | &app.mcp.db.lock().unwrap(), |
| 722 | &format!("pending:{}", hash(id)), |
| 723 | )?; |
| 724 | if pending.is_null() { |
| 725 | return Err(Error::new( |
| 726 | 404, |
| 727 | "This connection request expired. Start it again.", |
| 728 | )); |
| 729 | } |
| 730 | if !pending["owner"].is_null() && pending["owner"] != owner_id { |
| 731 | return Err(Error::new( |
| 732 | 403, |
| 733 | "This connection request belongs to another account.", |
| 734 | )); |
| 735 | } |
| 736 | string(&pending["resource"]).to_owned() |
| 737 | } else { |
| 738 | String::new() |
| 739 | }; |
| 740 | let agent_consent = consent_resource == app.mcp.resource("agents"); |
| 741 | let shale_consent = consent_resource == app.mcp.resource("shale"); |
| 742 | let catalog = CATALOGS |
| 743 | .iter() |
| 744 | .find(|(id, _, _)| consent_resource == app.mcp.resource(id)) |
| 745 | .map(|(id, _, _)| *id); |
| 746 | let mut linked = true; |
| 747 | let mut resource_error = None; |
| 748 | let resources: Vec<Value> = if body["deny"] == true { |
| 749 | Vec::new() |
| 750 | } else if agent_consent { |
| 751 | relay::machines(&app.mcp.db.lock().unwrap(), owner_id)? |
| 752 | .into_iter() |
| 753 | .map(|m| json!({"id":m["id"],"name":m["name"]})) |
| 754 | .collect() |
| 755 | } else if shale_consent { |
| 756 | let available = if body["resources"] == "all" { |
| 757 | shale::verified_session(&app, owner_id) |
| 758 | .await |
| 759 | .map(|_| Vec::new()) |
| 760 | } else { |
| 761 | shale::repositories(&app, owner_id).await |
| 762 | }; |
| 763 | match available { |
| 764 | Ok(repositories) => repositories, |
| 765 | Err(error) if error.status == 401 => { |
| 766 | linked = false; |
| 767 | Vec::new() |
| 768 | } |
| 769 | Err(error) if method == Method::GET => { |
| 770 | resource_error = Some(error.message); |
| 771 | Vec::new() |
| 772 | } |
| 773 | Err(error) => return Err(error), |
| 774 | } |
| 775 | } else if consent_resource == app.mcp.resource("observability") |
| 776 | && array(&me["sections"]).iter().any(|s| s == "admin") |
| 777 | && array(&owner["groups"]) |
| 778 | .iter() |
| 779 | .any(|g| g["name"] == "infra-admin") |
| 780 | { |
| 781 | core::scan(app.clone()) |
| 782 | .await? |
| 783 | .value |
| 784 | .as_object() |
| 785 | .unwrap() |
| 786 | .keys() |
| 787 | .map(|id| json!({"id":id,"name":id})) |
| 788 | .collect() |
| 789 | } else { |
| 790 | Vec::new() |
| 791 | }; |
| 792 | let mut db = app.mcp.db.lock().unwrap(); |
| 793 | let tx = db.transaction()?; |
| 794 | let value = match parts { |
| 795 | [] if method == Method::GET => { |
| 796 | let mut grants = list(&tx, "grant:")?; |
| 797 | grants.retain(|g| g["user"] == owner_id); |
| 798 | let machines = relay::machines(&tx, owner_id)?; |
| 799 | let connections = grants.iter().map(|grant| { |
| 800 | let name = if grant["client"].is_null() {grant["name"].clone()} else {get(&tx, &format!("client:{}", string(&grant["client"])))?["client_name"].clone()}; |
| 801 | let resources = if grant["resource"] == app.mcp.resource("shale") && grant["resources"] == "all" {json!("all")} else {json!(array(&grant["resources"]).iter().map(|id| if grant["resource"] == app.mcp.resource("agents") {machines.iter().find(|m| m["id"] == *id).map(|m| m["name"].clone()).unwrap_or_else(|| json!("Unlinked machine"))} else {id.clone()}).collect::<Vec<_>>())}; |
| 802 | let catalog = CATALOGS.iter().find(|(id, _, _)| grant["resource"] == app.mcp.resource(id)).map(|(id, _, _)| *id); |
| 803 | Ok(json!({"id":grant["id"],"name":name,"catalog":catalog,"resources":resources,"scopes":grant["scopes"],"createdAt":grant["createdAt"]})) |
| 804 | }).collect::<Result<Vec<_>>>()?; |
| 805 | let shale = get(&tx, &format!("shale-session:{owner_id}"))?; |
| 806 | let catalogs: Vec<_> = CATALOGS |
| 807 | .iter() |
| 808 | .map(|(id, name, _)| json!({"id":id,"name":name,"endpoint":app.mcp.resource(id)})) |
| 809 | .collect(); |
| 810 | json!({"catalogs":catalogs,"connections":connections,"machines":app.relay.view(machines,None),"shale":if shale["origin"] != app.shale.origin.as_str() {Value::Null} else {json!({"linkedAt":shale["linkedAt"]})}}) |
| 811 | } |
| 812 | ["consent", id] if method == Method::GET || method == Method::POST => { |
| 813 | let key = format!("pending:{}", hash(id)); |
| 814 | let mut pending = get(&tx, &key)?; |
| 815 | if pending.is_null() { |
| 816 | return Err(Error::new( |
| 817 | 404, |
| 818 | "This connection request expired. Start it again.", |
| 819 | )); |
| 820 | } |
| 821 | if !pending["owner"].is_null() && pending["owner"] != owner_id { |
| 822 | return Err(Error::new( |
| 823 | 403, |
| 824 | "This connection request belongs to another account.", |
| 825 | )); |
| 826 | } |
| 827 | pending["owner"] = json!(owner_id); |
| 828 | if method == Method::POST && body["deny"] == true { |
| 829 | delete(&tx, &key)?; |
| 830 | let mut target = redirect(string(&pending["redirect"]))?; |
| 831 | target |
| 832 | .query_pairs_mut() |
| 833 | .append_pair("error", "access_denied") |
| 834 | .append_pair("iss", app.mcp.origin.as_str()); |
| 835 | if let Some(state) = pending["state"].as_str() { |
| 836 | target.query_pairs_mut().append_pair("state", state); |
| 837 | } |
| 838 | tx.commit()?; |
| 839 | return Ok(axum::Json(json!({"redirect":target.as_str()})).into_response()); |
| 840 | } |
| 841 | if method == Method::GET { |
| 842 | tx.execute( |
| 843 | "UPDATE records SET value=? WHERE key=?", |
| 844 | rusqlite::params![pending.to_string(), key], |
| 845 | )?; |
| 846 | json!({"client":get(&tx,&format!("client:{}",string(&pending["client"])))?["client_name"],"catalog":catalog,"account":owner["username"],"redirectHost":redirect(string(&pending["redirect"]))?.host_str(),"scopes":pending["scopes"],"resources":resources,"linked":linked,"resourceError":resource_error}) |
| 847 | } else { |
| 848 | if shale_consent |
| 849 | && (!linked |
| 850 | || get(&tx, &format!("shale-session:{owner_id}"))?["origin"] |
| 851 | != app.shale.origin.as_str()) |
| 852 | { |
| 853 | return Err(Error::new( |
| 854 | 401, |
| 855 | "Link your Shale account before allowing repository access.", |
| 856 | )); |
| 857 | } |
| 858 | let chosen = chosen_resources(&body, &resources, shale_consent)?; |
| 859 | if list(&tx, "grant:")? |
| 860 | .iter() |
| 861 | .filter(|g| g["user"] == owner_id) |
| 862 | .count() |
| 863 | >= 256 |
| 864 | { |
| 865 | return Err(Error::new( |
| 866 | 409, |
| 867 | "Remove an unused connection before adding another.", |
| 868 | )); |
| 869 | } |
| 870 | let grant_id = uuid::Uuid::new_v4().to_string(); |
| 871 | let mut grant = json!({"id":grant_id,"user":owner_id,"client":pending["client"],"resource":pending["resource"],"scopes":pending["scopes"],"resources":chosen,"createdAt":now()}); |
| 872 | if agent_consent { |
| 873 | grant["targets"] = json!(chosen); |
| 874 | } |
| 875 | put(&tx, &format!("grant:{grant_id}"), &grant, 0)?; |
| 876 | let code = secret(); |
| 877 | pending["grant"] = json!(grant_id); |
| 878 | put(&tx, &format!("code:{}", hash(&code)), &pending, 300)?; |
| 879 | delete(&tx, &key)?; |
| 880 | let mut target = redirect(string(&pending["redirect"]))?; |
| 881 | target |
| 882 | .query_pairs_mut() |
| 883 | .append_pair("code", &code) |
| 884 | .append_pair("iss", app.mcp.origin.as_str()); |
| 885 | if let Some(state) = pending["state"].as_str() { |
| 886 | target.query_pairs_mut().append_pair("state", state); |
| 887 | } |
| 888 | json!({"redirect":target.as_str()}) |
| 889 | } |
| 890 | } |
| 891 | ["relay", rest @ ..] => relay::manage(&app, &tx, rest, method, owner_id, &body)?, |
| 892 | ["connections", id] |
| 893 | if method == Method::GET || method == Method::POST || method == Method::DELETE => |
| 894 | { |
| 895 | let key = format!("grant:{id}"); |
| 896 | let mut grant = get(&tx, &key)?; |
| 897 | if grant["user"] != owner_id { |
| 898 | return Err(Error::new(404, "No connection with that ID.")); |
| 899 | } |
| 900 | if method == Method::DELETE { |
| 901 | revoke(&tx, id)?; |
| 902 | Value::Null |
| 903 | } else if method == Method::GET { |
| 904 | json!({"resources": resources, "selected": grant["resources"], "linked": linked,"resourceError":resource_error}) |
| 905 | } else { |
| 906 | if shale_consent && !linked { |
| 907 | return Err(Error::new( |
| 908 | 401, |
| 909 | "Link your Shale account before allowing repository access.", |
| 910 | )); |
| 911 | } |
| 912 | let chosen = chosen_resources(&body, &resources, shale_consent)?; |
| 913 | grant["resources"] = json!(chosen); |
| 914 | if agent_consent { |
| 915 | grant["targets"] = json!(chosen); |
| 916 | } |
| 917 | tx.execute( |
| 918 | "UPDATE records SET value=? WHERE key=?", |
| 919 | rusqlite::params![grant.to_string(), key], |
| 920 | )?; |
| 921 | Value::Null |
| 922 | } |
| 923 | } |
| 924 | _ => return Err(Error::new(404, "No endpoint here.")), |
| 925 | }; |
| 926 | tx.commit()?; |
| 927 | if matches!(parts, ["relay", "pair"] | ["relay", "machines", _]) { |
| 928 | app.relay.changes.send_replace(()); |
| 929 | } |
| 930 | Ok(if value.is_null() { |
| 931 | StatusCode::NO_CONTENT.into_response() |
| 932 | } else { |
| 933 | axum::Json(value).into_response() |
| 934 | }) |
| 935 | } |
| 936 | |
| 937 | #[cfg(test)] |
| 938 | mod tests { |
| 939 | #[test] |
| 940 | fn resource_updates_reject_empty_duplicates_and_foreign_choices() { |
| 941 | let resources = vec![json!({"id":"alpha"}), json!({"id":"beta"})]; |
| 942 | assert_eq!( |
| 943 | chosen_resources(&json!({"resources":["beta"]}), &resources, false).unwrap(), |
| 944 | json!(["beta"]) |
| 945 | ); |
| 946 | for selected in [ |
| 947 | json!([]), |
| 948 | json!(["alpha", "alpha"]), |
| 949 | json!(["foreign"]), |
| 950 | json!([null]), |
| 951 | ] { |
| 952 | assert!(chosen_resources(&json!({"resources":selected}), &resources, false).is_err()); |
| 953 | } |
| 954 | assert_eq!( |
| 955 | chosen_resources(&json!({"resources":"all"}), &[], true).unwrap(), |
| 956 | json!("all") |
| 957 | ); |
| 958 | assert!(chosen_resources(&json!({"resources":"all"}), &resources, false).is_err()); |
| 959 | } |
| 960 | |
| 961 | use super::*; |
| 962 | struct Fixture { |
| 963 | store: Arc<Store>, |
| 964 | path: std::path::PathBuf, |
| 965 | } |
| 966 | impl Fixture { |
| 967 | fn new() -> Self { |
| 968 | let path = std::env::temp_dir() |
| 969 | .canonicalize() |
| 970 | .unwrap() |
| 971 | .join(format!("studio-mcp-test-{}", uuid::Uuid::new_v4())); |
| 972 | let store = Arc::new(Store::new(&path, "https://globe.studio.test").unwrap()); |
| 973 | Self { store, path } |
| 974 | } |
| 975 | fn code(&self, client: &str, code: &str) -> HashMap<String, String> { |
| 976 | let grant = uuid::Uuid::new_v4().to_string(); |
| 977 | let db = self.store.db.lock().unwrap(); |
| 978 | put(&db, &format!("client:{client}"), &json!({"client_id":client,"token_endpoint_auth_method":"none","redirect_uris":["http://127.0.0.1:20001/callback"]}),0).unwrap(); |
| 979 | put(&db, &format!("grant:{grant}"), &json!({"id":grant,"user":"one","client":client,"resource":self.store.resource("observability"),"scopes":["observability:read","offline_access"],"resources":["allowed"]}),0).unwrap(); |
| 980 | let verifier = "v".repeat(43); |
| 981 | put(&db, &format!("code:{}",hash(code)), &json!({"client":client,"redirect":"http://127.0.0.1:20001/callback","challenge":URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes())),"resource":self.store.resource("observability"),"grant":grant}),300).unwrap(); |
| 982 | fields( |
| 983 | json!({"grant_type":"authorization_code","client_id":client,"code":code,"code_verifier":verifier,"redirect_uri":"http://127.0.0.1:20001/callback","resource":self.store.resource("observability")}), |
| 984 | ) |
| 985 | } |
| 986 | } |
| 987 | impl Drop for Fixture { |
| 988 | fn drop(&mut self) { |
| 989 | std::fs::remove_dir_all(&self.path).unwrap(); |
| 990 | } |
| 991 | } |
| 992 | fn fields(value: Value) -> HashMap<String, String> { |
| 993 | value |
| 994 | .as_object() |
| 995 | .unwrap() |
| 996 | .iter() |
| 997 | .map(|(k, v)| (k.clone(), string(v).to_owned())) |
| 998 | .collect() |
| 999 | } |
| 1000 | async fn value(response: Response) -> Value { |
| 1001 | serde_json::from_slice( |
| 1002 | &axum::body::to_bytes(response.into_body(), 65536) |
| 1003 | .await |
| 1004 | .unwrap(), |
| 1005 | ) |
| 1006 | .unwrap() |
| 1007 | } |
| 1008 | fn bearer(token: &str) -> HeaderMap { |
| 1009 | let mut headers = HeaderMap::new(); |
| 1010 | headers.insert("authorization", format!("Bearer {token}").parse().unwrap()); |
| 1011 | headers |
| 1012 | } |
| 1013 | #[test] |
| 1014 | fn revocation_removes_durable_keys_and_family_records() { |
| 1015 | let fixture = Fixture::new(); |
| 1016 | let db = fixture.store.db.lock().unwrap(); |
| 1017 | put(&db, "grant:revoked", &json!({"id":"revoked"}), 0).unwrap(); |
| 1018 | for kind in ["access", "refresh", "used", "code"] { |
| 1019 | put(&db, &format!("{kind}:one"), &json!({"grant":"revoked"}), 0).unwrap(); |
| 1020 | put( |
| 1021 | &db, |
| 1022 | &format!("{kind}:other"), |
| 1023 | &json!({"grant":"retained"}), |
| 1024 | 0, |
| 1025 | ) |
| 1026 | .unwrap(); |
| 1027 | } |
| 1028 | revoke(&db, "revoked").unwrap(); |
| 1029 | assert!(get(&db, "grant:revoked").unwrap().is_null()); |
| 1030 | for kind in ["access", "refresh", "used", "code"] { |
| 1031 | assert!(get(&db, &format!("{kind}:one")).unwrap().is_null()); |
| 1032 | assert!(!get(&db, &format!("{kind}:other")).unwrap().is_null()); |
| 1033 | } |
| 1034 | } |
| 1035 | #[tokio::test] |
| 1036 | async fn code_binds_pkce_client_redirect_and_audience_without_consuming_on_failure() { |
| 1037 | let fixture = Fixture::new(); |
| 1038 | let input = fixture.code("one", "owned-code"); |
| 1039 | fixture.code("two", "other-code"); |
| 1040 | for (key, wrong) in [ |
| 1041 | ("client_id", "two"), |
| 1042 | ("code_verifier", "wrong"), |
| 1043 | ("redirect_uri", "http://127.0.0.1:20002/callback"), |
| 1044 | ("resource", "https://other.invalid/mcp/observability"), |
| 1045 | ] { |
| 1046 | let mut attempt = input.clone(); |
| 1047 | attempt.insert(key.into(), wrong.into()); |
| 1048 | assert!( |
| 1049 | fixture |
| 1050 | .store |
| 1051 | .oauth( |
| 1052 | "token", |
| 1053 | &Method::POST, |
| 1054 | &attempt, |
| 1055 | &HeaderMap::new(), |
| 1056 | Value::Null |
| 1057 | ) |
| 1058 | .is_err() |
| 1059 | ); |
| 1060 | } |
| 1061 | let tokens = value( |
| 1062 | fixture |
| 1063 | .store |
| 1064 | .oauth( |
| 1065 | "token", |
| 1066 | &Method::POST, |
| 1067 | &input, |
| 1068 | &HeaderMap::new(), |
| 1069 | Value::Null, |
| 1070 | ) |
| 1071 | .unwrap(), |
| 1072 | ) |
| 1073 | .await; |
| 1074 | assert!( |
| 1075 | fixture |
| 1076 | .store |
| 1077 | .oauth( |
| 1078 | "token", |
| 1079 | &Method::POST, |
| 1080 | &input, |
| 1081 | &HeaderMap::new(), |
| 1082 | Value::Null |
| 1083 | ) |
| 1084 | .is_err() |
| 1085 | ); |
| 1086 | assert!( |
| 1087 | fixture |
| 1088 | .store |
| 1089 | .authenticate( |
| 1090 | &bearer(string(&tokens["access_token"])), |
| 1091 | &fixture.store.resource("observability") |
| 1092 | ) |
| 1093 | .is_ok() |
| 1094 | ); |
| 1095 | assert!( |
| 1096 | fixture |
| 1097 | .store |
| 1098 | .authenticate( |
| 1099 | &bearer(string(&tokens["access_token"])), |
| 1100 | "https://other.invalid/mcp/observability" |
| 1101 | ) |
| 1102 | .is_err() |
| 1103 | ); |
| 1104 | assert_eq!( |
| 1105 | std::fs::metadata(fixture.path.join("connections.sqlite")) |
| 1106 | .unwrap() |
| 1107 | .permissions() |
| 1108 | .mode() |
| 1109 | & 0o777, |
| 1110 | 0o600 |
| 1111 | ); |
| 1112 | let rows = fixture |
| 1113 | .store |
| 1114 | .db |
| 1115 | .lock() |
| 1116 | .unwrap() |
| 1117 | .prepare("SELECT key,value FROM records") |
| 1118 | .unwrap() |
| 1119 | .query_map([], |row| { |
| 1120 | Ok(format!( |
| 1121 | "{} {}", |
| 1122 | row.get::<_, String>(0)?, |
| 1123 | row.get::<_, String>(1)? |
| 1124 | )) |
| 1125 | }) |
| 1126 | .unwrap() |
| 1127 | .collect::<std::result::Result<Vec<_>, _>>() |
| 1128 | .unwrap() |
| 1129 | .join("\n"); |
| 1130 | for token in [ |
| 1131 | "owned-code", |
| 1132 | string(&tokens["access_token"]), |
| 1133 | string(&tokens["refresh_token"]), |
| 1134 | ] { |
| 1135 | assert!(!rows.contains(token)); |
| 1136 | } |
| 1137 | } |
| 1138 | #[tokio::test] |
| 1139 | async fn refresh_replay_revokes_family_but_another_client_cannot_revoke_it() { |
| 1140 | let fixture = Fixture::new(); |
| 1141 | let input = fixture.code("one", "owned-code"); |
| 1142 | fixture.code("two", "other-code"); |
| 1143 | let tokens = value( |
| 1144 | fixture |
| 1145 | .store |
| 1146 | .oauth( |
| 1147 | "token", |
| 1148 | &Method::POST, |
| 1149 | &input, |
| 1150 | &HeaderMap::new(), |
| 1151 | Value::Null, |
| 1152 | ) |
| 1153 | .unwrap(), |
| 1154 | ) |
| 1155 | .await; |
| 1156 | let refresh = fields( |
| 1157 | json!({"grant_type":"refresh_token","client_id":"one","refresh_token":tokens["refresh_token"]}), |
| 1158 | ); |
| 1159 | let mut wrong = refresh.clone(); |
| 1160 | wrong.insert("client_id".into(), "two".into()); |
| 1161 | assert!( |
| 1162 | fixture |
| 1163 | .store |
| 1164 | .oauth( |
| 1165 | "token", |
| 1166 | &Method::POST, |
| 1167 | &wrong, |
| 1168 | &HeaderMap::new(), |
| 1169 | Value::Null |
| 1170 | ) |
| 1171 | .is_err() |
| 1172 | ); |
| 1173 | fixture |
| 1174 | .store |
| 1175 | .oauth( |
| 1176 | "revoke", |
| 1177 | &Method::POST, |
| 1178 | &fields(json!({"client_id":"two","token":tokens["access_token"]})), |
| 1179 | &HeaderMap::new(), |
| 1180 | Value::Null, |
| 1181 | ) |
| 1182 | .unwrap(); |
| 1183 | let access = bearer(string(&tokens["access_token"])); |
| 1184 | assert!( |
| 1185 | fixture |
| 1186 | .store |
| 1187 | .authenticate(&access, &fixture.store.resource("observability")) |
| 1188 | .is_ok() |
| 1189 | ); |
| 1190 | let rotated = value( |
| 1191 | fixture |
| 1192 | .store |
| 1193 | .oauth( |
| 1194 | "token", |
| 1195 | &Method::POST, |
| 1196 | &refresh, |
| 1197 | &HeaderMap::new(), |
| 1198 | Value::Null, |
| 1199 | ) |
| 1200 | .unwrap(), |
| 1201 | ) |
| 1202 | .await; |
| 1203 | assert_ne!(rotated["refresh_token"], tokens["refresh_token"]); |
| 1204 | assert!( |
| 1205 | fixture |
| 1206 | .store |
| 1207 | .oauth( |
| 1208 | "token", |
| 1209 | &Method::POST, |
| 1210 | &wrong, |
| 1211 | &HeaderMap::new(), |
| 1212 | Value::Null |
| 1213 | ) |
| 1214 | .is_err() |
| 1215 | ); |
| 1216 | assert!( |
| 1217 | fixture |
| 1218 | .store |
| 1219 | .authenticate(&access, &fixture.store.resource("observability")) |
| 1220 | .is_ok() |
| 1221 | ); |
| 1222 | assert!( |
| 1223 | fixture |
| 1224 | .store |
| 1225 | .oauth( |
| 1226 | "token", |
| 1227 | &Method::POST, |
| 1228 | &refresh, |
| 1229 | &HeaderMap::new(), |
| 1230 | Value::Null |
| 1231 | ) |
| 1232 | .is_err() |
| 1233 | ); |
| 1234 | assert!( |
| 1235 | fixture |
| 1236 | .store |
| 1237 | .authenticate(&access, &fixture.store.resource("observability")) |
| 1238 | .is_err() |
| 1239 | ); |
| 1240 | assert!( |
| 1241 | fixture |
| 1242 | .store |
| 1243 | .authenticate( |
| 1244 | &bearer(string(&rotated["access_token"])), |
| 1245 | &fixture.store.resource("observability") |
| 1246 | ) |
| 1247 | .is_err() |
| 1248 | ); |
| 1249 | assert!(fixture.store.oauth("token",&Method::POST,&fields(json!({"grant_type":"refresh_token","client_id":"one","refresh_token":rotated["refresh_token"]})),&HeaderMap::new(),Value::Null).is_err()); |
| 1250 | } |
| 1251 | #[test] |
| 1252 | fn concurrent_code_exchange_has_one_winner() { |
| 1253 | let fixture = Fixture::new(); |
| 1254 | let input = fixture.code("one", "concurrent-code"); |
| 1255 | let start = Arc::new(std::sync::Barrier::new(9)); |
| 1256 | let workers = (0..8) |
| 1257 | .map(|_| { |
| 1258 | let store = fixture.store.clone(); |
| 1259 | let input = input.clone(); |
| 1260 | let start = start.clone(); |
| 1261 | std::thread::spawn(move || { |
| 1262 | start.wait(); |
| 1263 | store |
| 1264 | .oauth( |
| 1265 | "token", |
| 1266 | &Method::POST, |
| 1267 | &input, |
| 1268 | &HeaderMap::new(), |
| 1269 | Value::Null, |
| 1270 | ) |
| 1271 | .is_ok() |
| 1272 | }) |
| 1273 | }) |
| 1274 | .collect::<Vec<_>>(); |
| 1275 | start.wait(); |
| 1276 | assert_eq!( |
| 1277 | workers |
| 1278 | .into_iter() |
| 1279 | .map(|w| w.join().unwrap() as u32) |
| 1280 | .sum::<u32>(), |
| 1281 | 1 |
| 1282 | ); |
| 1283 | } |
| 1284 | #[tokio::test] |
| 1285 | async fn confidential_client_secret_is_required_and_hashed() { |
| 1286 | let fixture = Fixture::new(); |
| 1287 | let client=value(fixture.store.oauth("register",&Method::POST,&HashMap::new(),&HeaderMap::new(),json!({"client_name":"Confidential","redirect_uris":["https://client.invalid/callback"],"token_endpoint_auth_method":"client_secret_basic"})).unwrap()).await; |
| 1288 | let id = string(&client["client_id"]); |
| 1289 | let stored = get(&fixture.store.db.lock().unwrap(), &format!("client:{id}")).unwrap(); |
| 1290 | assert!(stored.get("client_secret").is_none()); |
| 1291 | assert_ne!(stored["secret_hash"], client["client_secret"]); |
| 1292 | let input = fields(json!({"client_id":id,"token":"unknown"})); |
| 1293 | assert!( |
| 1294 | fixture |
| 1295 | .store |
| 1296 | .oauth( |
| 1297 | "revoke", |
| 1298 | &Method::POST, |
| 1299 | &input, |
| 1300 | &HeaderMap::new(), |
| 1301 | Value::Null |
| 1302 | ) |
| 1303 | .is_err() |
| 1304 | ); |
| 1305 | let mut headers = HeaderMap::new(); |
| 1306 | headers.insert( |
| 1307 | "authorization", |
| 1308 | format!( |
| 1309 | "Basic {}", |
| 1310 | STANDARD.encode(format!("{id}:{}", string(&client["client_secret"]))) |
| 1311 | ) |
| 1312 | .parse() |
| 1313 | .unwrap(), |
| 1314 | ); |
| 1315 | assert!( |
| 1316 | fixture |
| 1317 | .store |
| 1318 | .oauth("revoke", &Method::POST, &input, &headers, Value::Null) |
| 1319 | .is_ok() |
| 1320 | ); |
| 1321 | let mut duplicate = input.clone(); |
| 1322 | duplicate.insert( |
| 1323 | "client_secret".into(), |
| 1324 | string(&client["client_secret"]).into(), |
| 1325 | ); |
| 1326 | assert!( |
| 1327 | fixture |
| 1328 | .store |
| 1329 | .oauth("revoke", &Method::POST, &duplicate, &headers, Value::Null) |
| 1330 | .is_err() |
| 1331 | ); |
| 1332 | } |
| 1333 | } |