1use crate::*;
2use base64::{
3 Engine,
4 engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD},
5};
6use rusqlite::{Connection, OpenFlags, OptionalExtension};
7use sha2::{Digest, Sha256};
8use std::os::unix::fs::{DirBuilderExt, PermissionsExt};
9
10const 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
28pub struct Store {
29 pub(crate) db: Mutex<Connection>,
30 pub origin: url::Url,
31}
32pub(crate) fn secret() -> String {
33 URL_SAFE_NO_PAD.encode(rand::random::<[u8; 32]>())
34}
35pub(crate) fn hash(value: &str) -> String {
36 format!("{:x}", Sha256::digest(value.as_bytes()))
37}
38pub(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}
51pub(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}
73pub(crate) fn delete(db: &Connection, key: &str) -> Result<()> {
74 db.execute("DELETE FROM records WHERE key=?", [key])?;
75 Ok(())
76}
77pub(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}
84pub(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}
93fn fail(code: &str) -> Error {
94 Error::new(400, code)
95}
96fn 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}
110impl 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!("/mcp?request={}", 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)]
492pub(crate) struct Grant(pub Value);
493pub fn router<H: rmcp::ServerHandler>(
494 app: Arc<App>,
495 catalog: &str,
496 handler: impl Fn() -> std::result::Result<H, std::io::Error> + Send + Sync + 'static,
497) -> Router {
498 use rmcp::transport::streamable_http_server::{
499 StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager,
500 };
501 let resource = app.mcp.resource(catalog);
502 let metadata = format!(
503 "{}.well-known/oauth-protected-resource/mcp/{catalog}",
504 app.mcp.origin.as_str()
505 );
506 let mut config = StreamableHttpServerConfig::default();
507 config.legacy_session_mode = false;
508 config.json_response = true;
509 config.allowed_hosts =
510 vec![app.mcp.origin[url::Position::BeforeHost..url::Position::AfterPort].to_owned()];
511 config.allowed_origins = vec![format!(
512 "{}://{}:{}",
513 app.mcp.origin.scheme(),
514 app.mcp.origin.host_str().unwrap(),
515 app.mcp.origin.port_or_known_default().unwrap()
516 )];
517 let service =
518 StreamableHttpService::new(handler, Arc::new(LocalSessionManager::default()), config);
519 Router::new()
520 .nest_service(&format!("/mcp/{catalog}"), service)
521 .layer(axum::middleware::from_fn(
522 move |mut request: Request, next: axum::middleware::Next| {
523 let (app, resource, metadata) = (app.clone(), resource.clone(), metadata.clone());
524 async move {
525 if request.headers().get("origin").is_some_and(|origin| {
526 origin.to_str().ok()
527 != Some(app.mcp.origin.origin().ascii_serialization().as_str())
528 }) {
529 return Error::new(403, "This origin cannot use the connector.")
530 .into_response();
531 }
532 match app.mcp.authenticate(request.headers(), &resource) {
533 Ok(grant) => {
534 request.extensions_mut().insert(Grant(grant));
535 if let Some(value) = request.headers_mut().get_mut("authorization") {
536 value.set_sensitive(true);
537 }
538 next.run(request).await
539 }
540 Err(_) => (
541 StatusCode::UNAUTHORIZED,
542 [(
543 "www-authenticate",
544 format!("Bearer resource_metadata=\"{metadata}\""),
545 )],
546 "This connection expired. Connect again.",
547 )
548 .into_response(),
549 }
550 }
551 },
552 ))
553}
554
555pub fn public(path: &str) -> bool {
556 path.starts_with("/oauth/")
557 || path.starts_with("/mcp/")
558 || path.starts_with("/.well-known/oauth-")
559 || path == "/pairing"
560 || path == "/agent/connect"
561 || path.starts_with("/api/v1/")
562}
563pub async fn oauth(State(app): State<Arc<App>>, request: Request) -> Response {
564 if request.uri().path().starts_with("/oauth/shale/") {
565 return shale::oauth(app, request).await;
566 }
567 let result = async {
568 let path = request.uri().path().to_owned();
569 if path=="/.well-known/oauth-authorization-server" && request.method()==Method::GET {
570 let base = app.mcp.origin.as_str();
571 let scopes: std::collections::BTreeSet<_> = CATALOGS.iter().flat_map(|(_, _, scopes)| scopes.iter()).collect();
572 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());
573 }
574 if request.method() == Method::GET {
575 for (catalog, _, scopes) in CATALOGS {
576 if path == format!("/.well-known/oauth-protected-resource/mcp/{catalog}") {
577 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());
578 }
579 }
580 }
581 let method = request.method().clone();
582 let headers = request.headers().clone();
583 let query = request.uri().query().unwrap_or_default().to_owned();
584 let bytes = axum::body::to_bytes(request.into_body(),65536).await.map_err(|_| fail("invalid_request"))?;
585 let registration = path=="/oauth/register";
586 let encoded = if method==Method::GET {query.as_bytes()} else {&bytes};
587 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")); }
588 let mut input = HashMap::new();
589 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"));} } }
590 let body = if registration {serde_json::from_slice(&bytes).map_err(|_| fail("invalid_client_metadata"))?} else {Value::Null};
591 if path == "/oauth/token" && method == Method::POST {
592 let (_, grant) = app.mcp.exchange(&app.mcp.db.lock().unwrap(), &input, &headers)?;
593 let id = string(&grant["user"]);
594 let (profile, roles) = match tokio::try_join!(
595 host::call(json!({"operation":"iam.request","path":format!("/users/{id}"),"method":"GET","body":null})),
596 host::call(json!({"operation":"iam.request","path":format!("/users/{id}/role-mappings/realm"),"method":"GET","body":null}))
597 ) {
598 Ok(identity) => identity,
599 Err(error) if error.status == 404 => {
600 revoke(&app.mcp.db.lock().unwrap(), string(&grant["id"]))?;
601 return Err(fail("invalid_grant"));
602 }
603 Err(error) => return Err(error),
604 };
605 if profile["body"]["enabled"] != true || grant["resource"] == app.mcp.resource("observability") && !array(&roles["body"]).iter().any(|role| role["name"] == "infra-admin") {
606 revoke(&app.mcp.db.lock().unwrap(), string(&grant["id"]))?;
607 return Err(fail("invalid_grant"));
608 }
609 }
610 app.mcp.oauth(path.trim_start_matches("/oauth/"),&method,&input,&headers,body)
611 }.await;
612 let mut response = match result {
613 Ok(response) => response,
614 Err(error) => (
615 StatusCode::from_u16(if error.message == "invalid_client" {
616 401
617 } else {
618 error.status
619 })
620 .unwrap_or(StatusCode::BAD_REQUEST),
621 axum::Json(
622 json!({"error":if error.status >= 500 {"server_error"} else {&error.message}}),
623 ),
624 )
625 .into_response(),
626 };
627 response
628 .headers_mut()
629 .insert("cache-control", "no-store".parse().unwrap());
630 response
631 .headers_mut()
632 .insert("pragma", "no-cache".parse().unwrap());
633 response
634}
635
636pub async fn manage(
637 app: Arc<App>,
638 method: &Method,
639 parts: &[&str],
640 me: &Value,
641 body: Value,
642 headers: &HeaderMap,
643) -> Result<Response> {
644 if method != Method::GET
645 && (me["viewing"] == true
646 || headers.get("origin").and_then(|h| h.to_str().ok())
647 != Some(app.mcp.origin.origin().ascii_serialization().as_str()))
648 {
649 return Err(Error::new(
650 403,
651 "Open MCP settings from your own signed-in account.",
652 ));
653 }
654 let owner = users::self_user(&app, me).await?;
655 if owner["enabled"] != true {
656 return Err(Error::new(403, "This account is disabled."));
657 }
658 let owner_id = string(&owner["id"]);
659 if parts == ["shale"] {
660 return shale::manage(app, method, &owner, &body).await;
661 }
662 if parts == ["relay", "live"] && method == Method::GET {
663 let owner = owner_id.to_owned();
664 let stream = WatchStream::new(app.relay.changes.subscribe()).map(move |_| {
665 let machines = relay::machines(&app.mcp.db.lock().unwrap(), &owner)
666 .map(|machines| app.relay.view(machines, None));
667 machines
668 .map(|machines| Event::default().json_data(machines).unwrap())
669 .map_err(|error| std::io::Error::other(error.message))
670 });
671 return Ok(Sse::new(stream)
672 .keep_alive(axum::response::sse::KeepAlive::default())
673 .into_response());
674 }
675 let consent_resource = if let ["consent", id] = parts {
676 get(
677 &app.mcp.db.lock().unwrap(),
678 &format!("pending:{}", hash(id)),
679 )?["resource"]
680 .as_str()
681 .unwrap_or_default()
682 .to_owned()
683 } else {
684 String::new()
685 };
686 let agent_consent = consent_resource == app.mcp.resource("agents");
687 let shale_consent = consent_resource == app.mcp.resource("shale");
688 let mut linked = true;
689 let resources: Vec<Value> = if agent_consent {
690 relay::machines(&app.mcp.db.lock().unwrap(), owner_id)?
691 .into_iter()
692 .map(|m| json!({"id":m["id"],"name":m["name"]}))
693 .collect()
694 } else if shale_consent && body["deny"] != true {
695 match shale::repositories(&app, owner_id).await {
696 Ok(repositories) => repositories,
697 Err(error) if error.status == 401 => {
698 linked = false;
699 Vec::new()
700 }
701 Err(error) => return Err(error),
702 }
703 } else if consent_resource == app.mcp.resource("observability")
704 && array(&me["sections"]).iter().any(|s| s == "admin")
705 && array(&owner["groups"])
706 .iter()
707 .any(|g| g["name"] == "infra-admin")
708 {
709 core::scan(app.clone())
710 .await?
711 .value
712 .as_object()
713 .unwrap()
714 .keys()
715 .map(|id| json!({"id":id,"name":id}))
716 .collect()
717 } else {
718 Vec::new()
719 };
720 let mut db = app.mcp.db.lock().unwrap();
721 let tx = db.transaction()?;
722 let value = match parts {
723 [] if method == Method::GET => {
724 let mut grants = list(&tx, "grant:")?;
725 grants.retain(|g| g["user"] == owner_id);
726 let machines = relay::machines(&tx, owner_id)?;
727 let connections = grants.iter().map(|grant| {
728 let name = if grant["client"].is_null() {grant["name"].clone()} else {get(&tx, &format!("client:{}", string(&grant["client"])))?["client_name"].clone()};
729 let resources: Vec<_> = 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();
730 Ok(json!({"id":grant["id"],"name":name,"resources":resources,"scopes":grant["scopes"],"createdAt":grant["createdAt"]}))
731 }).collect::<Result<Vec<_>>>()?;
732 let shale = get(&tx, &format!("shale-session:{owner_id}"))?;
733 let catalogs: Vec<_> = CATALOGS
734 .iter()
735 .map(|(id, name, _)| json!({"name":name,"endpoint":app.mcp.resource(id)}))
736 .collect();
737 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"]})}})
738 }
739 ["consent", id] if method == Method::GET || method == Method::POST => {
740 let key = format!("pending:{}", hash(id));
741 let mut pending = get(&tx, &key)?;
742 if pending.is_null() {
743 return Err(Error::new(
744 404,
745 "This connection request expired. Start it again.",
746 ));
747 }
748 if !pending["owner"].is_null() && pending["owner"] != owner_id {
749 return Err(Error::new(
750 403,
751 "This connection request belongs to another account.",
752 ));
753 }
754 pending["owner"] = json!(owner_id);
755 if method == Method::POST && body["deny"] == true {
756 delete(&tx, &key)?;
757 let mut target = redirect(string(&pending["redirect"]))?;
758 target
759 .query_pairs_mut()
760 .append_pair("error", "access_denied")
761 .append_pair("iss", app.mcp.origin.as_str());
762 if let Some(state) = pending["state"].as_str() {
763 target.query_pairs_mut().append_pair("state", state);
764 }
765 tx.commit()?;
766 return Ok(axum::Json(json!({"redirect":target.as_str()})).into_response());
767 }
768 if method == Method::GET {
769 tx.execute(
770 "UPDATE records SET value=? WHERE key=?",
771 rusqlite::params![pending.to_string(), key],
772 )?;
773 json!({"client":get(&tx,&format!("client:{}",string(&pending["client"])))?["client_name"],"scopes":pending["scopes"],"resources":resources,"linked":linked})
774 } else {
775 if shale_consent
776 && get(&tx, &format!("shale-session:{owner_id}"))?["origin"]
777 != app.shale.origin.as_str()
778 {
779 return Err(Error::new(
780 401,
781 "Link your Shale account before allowing repository access.",
782 ));
783 }
784 let chosen = body["resources"]
785 .as_array()
786 .filter(|a| !a.is_empty() && a.len() <= resources.len())
787 .ok_or_else(|| Error::new(400, "Choose each available resource once."))?;
788 if chosen.iter().enumerate().any(|(index, r)| {
789 !resources.iter().any(|id| r == &id["id"]) || chosen[..index].contains(r)
790 }) {
791 return Err(Error::new(
792 403,
793 "Choose resources available to your account.",
794 ));
795 }
796 if list(&tx, "grant:")?
797 .iter()
798 .filter(|g| g["user"] == owner_id)
799 .count()
800 >= 256
801 {
802 return Err(Error::new(
803 409,
804 "Remove an unused connection before adding another.",
805 ));
806 }
807 let grant_id = uuid::Uuid::new_v4().to_string();
808 let mut grant = json!({"id":grant_id,"user":owner_id,"client":pending["client"],"resource":pending["resource"],"scopes":pending["scopes"],"resources":chosen,"createdAt":now()});
809 if agent_consent {
810 grant["targets"] = json!(chosen);
811 }
812 put(&tx, &format!("grant:{grant_id}"), &grant, 0)?;
813 let code = secret();
814 pending["grant"] = json!(grant_id);
815 put(&tx, &format!("code:{}", hash(&code)), &pending, 300)?;
816 delete(&tx, &key)?;
817 let mut target = redirect(string(&pending["redirect"]))?;
818 target
819 .query_pairs_mut()
820 .append_pair("code", &code)
821 .append_pair("iss", app.mcp.origin.as_str());
822 if let Some(state) = pending["state"].as_str() {
823 target.query_pairs_mut().append_pair("state", state);
824 }
825 json!({"redirect":target.as_str()})
826 }
827 }
828 ["relay", rest @ ..] => relay::manage(&app, &tx, rest, method, owner_id, &body)?,
829 ["connections", id] if method == Method::DELETE => {
830 let grant = get(&tx, &format!("grant:{id}"))?;
831 if grant["user"] != owner_id {
832 return Err(Error::new(404, "No connection with that ID."));
833 }
834 revoke(&tx, id)?;
835 Value::Null
836 }
837 _ => return Err(Error::new(404, "No endpoint here.")),
838 };
839 tx.commit()?;
840 if matches!(parts, ["relay", "pair"] | ["relay", "machines", _]) {
841 app.relay.changes.send_replace(());
842 }
843 Ok(if value.is_null() {
844 StatusCode::NO_CONTENT.into_response()
845 } else {
846 axum::Json(value).into_response()
847 })
848}
849
850#[cfg(test)]
851mod tests {
852 use super::*;
853 struct Fixture {
854 store: Arc<Store>,
855 path: std::path::PathBuf,
856 }
857 impl Fixture {
858 fn new() -> Self {
859 let path = std::env::temp_dir()
860 .canonicalize()
861 .unwrap()
862 .join(format!("studio-mcp-test-{}", uuid::Uuid::new_v4()));
863 let store = Arc::new(Store::new(&path, "https://globe.studio.test").unwrap());
864 Self { store, path }
865 }
866 fn code(&self, client: &str, code: &str) -> HashMap<String, String> {
867 let grant = uuid::Uuid::new_v4().to_string();
868 let db = self.store.db.lock().unwrap();
869 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();
870 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();
871 let verifier = "v".repeat(43);
872 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();
873 fields(
874 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")}),
875 )
876 }
877 }
878 impl Drop for Fixture {
879 fn drop(&mut self) {
880 std::fs::remove_dir_all(&self.path).unwrap();
881 }
882 }
883 fn fields(value: Value) -> HashMap<String, String> {
884 value
885 .as_object()
886 .unwrap()
887 .iter()
888 .map(|(k, v)| (k.clone(), string(v).to_owned()))
889 .collect()
890 }
891 async fn value(response: Response) -> Value {
892 serde_json::from_slice(
893 &axum::body::to_bytes(response.into_body(), 65536)
894 .await
895 .unwrap(),
896 )
897 .unwrap()
898 }
899 fn bearer(token: &str) -> HeaderMap {
900 let mut headers = HeaderMap::new();
901 headers.insert("authorization", format!("Bearer {token}").parse().unwrap());
902 headers
903 }
904 #[test]
905 fn revocation_removes_durable_keys_and_family_records() {
906 let fixture = Fixture::new();
907 let db = fixture.store.db.lock().unwrap();
908 put(&db, "grant:revoked", &json!({"id":"revoked"}), 0).unwrap();
909 for kind in ["access", "refresh", "used", "code"] {
910 put(&db, &format!("{kind}:one"), &json!({"grant":"revoked"}), 0).unwrap();
911 put(
912 &db,
913 &format!("{kind}:other"),
914 &json!({"grant":"retained"}),
915 0,
916 )
917 .unwrap();
918 }
919 revoke(&db, "revoked").unwrap();
920 assert!(get(&db, "grant:revoked").unwrap().is_null());
921 for kind in ["access", "refresh", "used", "code"] {
922 assert!(get(&db, &format!("{kind}:one")).unwrap().is_null());
923 assert!(!get(&db, &format!("{kind}:other")).unwrap().is_null());
924 }
925 }
926 #[tokio::test]
927 async fn code_binds_pkce_client_redirect_and_audience_without_consuming_on_failure() {
928 let fixture = Fixture::new();
929 let input = fixture.code("one", "owned-code");
930 fixture.code("two", "other-code");
931 for (key, wrong) in [
932 ("client_id", "two"),
933 ("code_verifier", "wrong"),
934 ("redirect_uri", "http://127.0.0.1:20002/callback"),
935 ("resource", "https://other.invalid/mcp/observability"),
936 ] {
937 let mut attempt = input.clone();
938 attempt.insert(key.into(), wrong.into());
939 assert!(
940 fixture
941 .store
942 .oauth(
943 "token",
944 &Method::POST,
945 &attempt,
946 &HeaderMap::new(),
947 Value::Null
948 )
949 .is_err()
950 );
951 }
952 let tokens = value(
953 fixture
954 .store
955 .oauth(
956 "token",
957 &Method::POST,
958 &input,
959 &HeaderMap::new(),
960 Value::Null,
961 )
962 .unwrap(),
963 )
964 .await;
965 assert!(
966 fixture
967 .store
968 .oauth(
969 "token",
970 &Method::POST,
971 &input,
972 &HeaderMap::new(),
973 Value::Null
974 )
975 .is_err()
976 );
977 assert!(
978 fixture
979 .store
980 .authenticate(
981 &bearer(string(&tokens["access_token"])),
982 &fixture.store.resource("observability")
983 )
984 .is_ok()
985 );
986 assert!(
987 fixture
988 .store
989 .authenticate(
990 &bearer(string(&tokens["access_token"])),
991 "https://other.invalid/mcp/observability"
992 )
993 .is_err()
994 );
995 assert_eq!(
996 std::fs::metadata(fixture.path.join("connections.sqlite"))
997 .unwrap()
998 .permissions()
999 .mode()
1000 & 0o777,
1001 0o600
1002 );
1003 let rows = fixture
1004 .store
1005 .db
1006 .lock()
1007 .unwrap()
1008 .prepare("SELECT key,value FROM records")
1009 .unwrap()
1010 .query_map([], |row| {
1011 Ok(format!(
1012 "{} {}",
1013 row.get::<_, String>(0)?,
1014 row.get::<_, String>(1)?
1015 ))
1016 })
1017 .unwrap()
1018 .collect::<std::result::Result<Vec<_>, _>>()
1019 .unwrap()
1020 .join("\n");
1021 for token in [
1022 "owned-code",
1023 string(&tokens["access_token"]),
1024 string(&tokens["refresh_token"]),
1025 ] {
1026 assert!(!rows.contains(token));
1027 }
1028 }
1029 #[tokio::test]
1030 async fn refresh_replay_revokes_family_but_another_client_cannot_revoke_it() {
1031 let fixture = Fixture::new();
1032 let input = fixture.code("one", "owned-code");
1033 fixture.code("two", "other-code");
1034 let tokens = value(
1035 fixture
1036 .store
1037 .oauth(
1038 "token",
1039 &Method::POST,
1040 &input,
1041 &HeaderMap::new(),
1042 Value::Null,
1043 )
1044 .unwrap(),
1045 )
1046 .await;
1047 let refresh = fields(
1048 json!({"grant_type":"refresh_token","client_id":"one","refresh_token":tokens["refresh_token"]}),
1049 );
1050 let mut wrong = refresh.clone();
1051 wrong.insert("client_id".into(), "two".into());
1052 assert!(
1053 fixture
1054 .store
1055 .oauth(
1056 "token",
1057 &Method::POST,
1058 &wrong,
1059 &HeaderMap::new(),
1060 Value::Null
1061 )
1062 .is_err()
1063 );
1064 fixture
1065 .store
1066 .oauth(
1067 "revoke",
1068 &Method::POST,
1069 &fields(json!({"client_id":"two","token":tokens["access_token"]})),
1070 &HeaderMap::new(),
1071 Value::Null,
1072 )
1073 .unwrap();
1074 let access = bearer(string(&tokens["access_token"]));
1075 assert!(
1076 fixture
1077 .store
1078 .authenticate(&access, &fixture.store.resource("observability"))
1079 .is_ok()
1080 );
1081 let rotated = value(
1082 fixture
1083 .store
1084 .oauth(
1085 "token",
1086 &Method::POST,
1087 &refresh,
1088 &HeaderMap::new(),
1089 Value::Null,
1090 )
1091 .unwrap(),
1092 )
1093 .await;
1094 assert_ne!(rotated["refresh_token"], tokens["refresh_token"]);
1095 assert!(
1096 fixture
1097 .store
1098 .oauth(
1099 "token",
1100 &Method::POST,
1101 &wrong,
1102 &HeaderMap::new(),
1103 Value::Null
1104 )
1105 .is_err()
1106 );
1107 assert!(
1108 fixture
1109 .store
1110 .authenticate(&access, &fixture.store.resource("observability"))
1111 .is_ok()
1112 );
1113 assert!(
1114 fixture
1115 .store
1116 .oauth(
1117 "token",
1118 &Method::POST,
1119 &refresh,
1120 &HeaderMap::new(),
1121 Value::Null
1122 )
1123 .is_err()
1124 );
1125 assert!(
1126 fixture
1127 .store
1128 .authenticate(&access, &fixture.store.resource("observability"))
1129 .is_err()
1130 );
1131 assert!(
1132 fixture
1133 .store
1134 .authenticate(
1135 &bearer(string(&rotated["access_token"])),
1136 &fixture.store.resource("observability")
1137 )
1138 .is_err()
1139 );
1140 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());
1141 }
1142 #[test]
1143 fn concurrent_code_exchange_has_one_winner() {
1144 let fixture = Fixture::new();
1145 let input = fixture.code("one", "concurrent-code");
1146 let start = Arc::new(std::sync::Barrier::new(9));
1147 let workers = (0..8)
1148 .map(|_| {
1149 let store = fixture.store.clone();
1150 let input = input.clone();
1151 let start = start.clone();
1152 std::thread::spawn(move || {
1153 start.wait();
1154 store
1155 .oauth(
1156 "token",
1157 &Method::POST,
1158 &input,
1159 &HeaderMap::new(),
1160 Value::Null,
1161 )
1162 .is_ok()
1163 })
1164 })
1165 .collect::<Vec<_>>();
1166 start.wait();
1167 assert_eq!(
1168 workers
1169 .into_iter()
1170 .map(|w| w.join().unwrap() as u32)
1171 .sum::<u32>(),
1172 1
1173 );
1174 }
1175 #[tokio::test]
1176 async fn confidential_client_secret_is_required_and_hashed() {
1177 let fixture = Fixture::new();
1178 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;
1179 let id = string(&client["client_id"]);
1180 let stored = get(&fixture.store.db.lock().unwrap(), &format!("client:{id}")).unwrap();
1181 assert!(stored.get("client_secret").is_none());
1182 assert_ne!(stored["secret_hash"], client["client_secret"]);
1183 let input = fields(json!({"client_id":id,"token":"unknown"}));
1184 assert!(
1185 fixture
1186 .store
1187 .oauth(
1188 "revoke",
1189 &Method::POST,
1190 &input,
1191 &HeaderMap::new(),
1192 Value::Null
1193 )
1194 .is_err()
1195 );
1196 let mut headers = HeaderMap::new();
1197 headers.insert(
1198 "authorization",
1199 format!(
1200 "Basic {}",
1201 STANDARD.encode(format!("{id}:{}", string(&client["client_secret"])))
1202 )
1203 .parse()
1204 .unwrap(),
1205 );
1206 assert!(
1207 fixture
1208 .store
1209 .oauth("revoke", &Method::POST, &input, &headers, Value::Null)
1210 .is_ok()
1211 );
1212 let mut duplicate = input.clone();
1213 duplicate.insert(
1214 "client_secret".into(),
1215 string(&client["client_secret"]).into(),
1216 );
1217 assert!(
1218 fixture
1219 .store
1220 .oauth("revoke", &Method::POST, &duplicate, &headers, Value::Null)
1221 .is_err()
1222 );
1223 }
1224}