| 1 | use crate::*; |
| 2 | use axum::body::Body; |
| 3 | use lol_html::{RewriteStrSettings, element, html_content::ContentType, text}; |
| 4 | |
| 5 | struct Profile { |
| 6 | name: String, |
| 7 | url: Option<String>, |
| 8 | icon: &'static str, |
| 9 | picture: Option<String>, |
| 10 | } |
| 11 | |
| 12 | fn profiles(auth: &auth::Store) -> Result<HashMap<String, Profile>> { |
| 13 | let db = auth.db.lock().unwrap(); |
| 14 | let mut query = db.prepare("SELECT json_extract(profile,'$.username'),coalesce(json_extract(profile,'$.firstName'),''),provider,subject,json_extract(profile,'$.attributes.picture[0]') FROM users JOIN external_identities ON user_id=users.id WHERE json_extract(profile,'$.kind')='guest'")?; |
| 15 | let rows = query.query_map([], |row| { |
| 16 | Ok(( |
| 17 | row.get::<_, String>(0)?, |
| 18 | row.get::<_, String>(1)?, |
| 19 | row.get::<_, String>(2)?, |
| 20 | row.get::<_, String>(3)?, |
| 21 | row.get::<_, Option<String>>(4)?, |
| 22 | )) |
| 23 | })?; |
| 24 | let mut profiles = HashMap::new(); |
| 25 | for row in rows { |
| 26 | let (username, name, provider, subject, picture) = row?; |
| 27 | if name.is_empty() { continue; } |
| 28 | let (url, icon) = match provider.as_str() { |
| 29 | "github" => { |
| 30 | if matches!(name.as_str(), "." | "..") { continue; } |
| 31 | let mut url = url::Url::parse("https://github.com/")?; |
| 32 | url.path_segments_mut().unwrap().pop_if_empty().push(&name); |
| 33 | (Some(String::from(url)), include_str!("../web/sso/github.svg")) |
| 34 | } |
| 35 | "astheno" => (None, include_str!("astheno.svg")), |
| 36 | _ => continue, |
| 37 | }; |
| 38 | profiles.insert( |
| 39 | username, |
| 40 | Profile { |
| 41 | name, |
| 42 | url, |
| 43 | icon, |
| 44 | picture: if provider == "github" { |
| 45 | Some(format!( |
| 46 | "https://avatars.githubusercontent.com/u/{subject}?s=64" |
| 47 | )) |
| 48 | } else { |
| 49 | picture |
| 50 | }, |
| 51 | }, |
| 52 | ); |
| 53 | } |
| 54 | Ok(profiles) |
| 55 | } |
| 56 | |
| 57 | fn account_path(path: &str) -> Option<&str> { |
| 58 | let name = path.strip_prefix("/~")?; |
| 59 | let name = name.strip_suffix('/').unwrap_or(name); |
| 60 | (name.starts_with("guest-") && !name.contains('/')).then_some(name) |
| 61 | } |
| 62 | |
| 63 | fn rewrite(html: &str, origin: &url::Url, profiles: &HashMap<String, Profile>) -> Result<String> { |
| 64 | use std::{cell::Cell, rc::Rc}; |
| 65 | let active = Rc::new(Cell::new(None::<bool>)); |
| 66 | let images = active.clone(); |
| 67 | let labels = active.clone(); |
| 68 | Ok(lol_html::rewrite_str(html, RewriteStrSettings::new() |
| 69 | .append_element_content_handler(element!("a[href]", |element| { |
| 70 | let profile = element.get_attribute("href") |
| 71 | .and_then(|href| origin.join(&href).ok()) |
| 72 | .filter(|target| target.origin() == origin.origin()) |
| 73 | .and_then(|target| account_path(target.path()).and_then(|name| profiles.get(name))); |
| 74 | active.set(profile.map(|profile| profile.picture.is_some())); |
| 75 | let Some(profile) = profile else { return Ok(()); }; |
| 76 | let closing = active.clone(); |
| 77 | element.on_end_tag(lol_html::end_tag!(move |_| { closing.set(None); Ok(()) }))?; |
| 78 | if let Some(url) = &profile.url { |
| 79 | element.set_attribute("href", url)?; |
| 80 | element.set_attribute("target", "_blank")?; |
| 81 | element.set_attribute("rel", "noreferrer")?; |
| 82 | } else { |
| 83 | for attribute in ["href", "target", "rel", "tabindex"] { element.remove_attribute(attribute); } |
| 84 | } |
| 85 | let style = element.get_attribute("style").unwrap_or_default(); |
| 86 | let cursor = if profile.url.is_some() { "" } else { ";cursor:default" }; |
| 87 | element.set_attribute("style", &format!("{style};text-decoration:none{cursor}"))?; |
| 88 | if let Some(picture) = &profile.picture { |
| 89 | let picture = picture.replace('&', "&amp;").replace('"', "&quot;").replace('<', "&lt;"); |
| 90 | element.prepend(&format!("<img src=\"{picture}\" alt=\"\" width=\"16\" height=\"16\" referrerpolicy=\"no-referrer\" style=\"width:1em;height:1em;object-fit:cover;vertical-align:-.125em;margin-right:.3em\">"), ContentType::Html); |
| 91 | } |
| 92 | let icon = profile.icon.trim().replace("<svg ", "<svg aria-hidden=\"true\" focusable=\"false\" style=\"color:inherit;stroke:currentColor;width:.85em;height:.85em;vertical-align:-.1em;margin-right:.05em\" "); |
| 93 | let icon = if profile.url.is_some() { icon.replace("<path ", "<path style=\"fill:currentColor;stroke:none\" ") } else { icon }; |
| 94 | element.append(&icon, ContentType::Html); |
| 95 | if profile.url.is_some() { |
| 96 | element.append("<span style=\"text-decoration:underline;text-decoration-color:currentColor\">", ContentType::Html); |
| 97 | } |
| 98 | element.append(&profile.name, ContentType::Text); |
| 99 | if profile.url.is_some() { element.append("</span>", ContentType::Html); } |
| 100 | Ok(()) |
| 101 | })) |
| 102 | .append_element_content_handler(element!("a[href] img, a[href] span", move |child| { |
| 103 | if let Some(replace_picture) = images.get() { |
| 104 | if child.tag_name() == "span" { child.remove_and_keep_content(); } |
| 105 | else if replace_picture { child.remove(); } |
| 106 | else { |
| 107 | child.set_attribute("alt", "")?; |
| 108 | child.set_attribute("referrerpolicy", "no-referrer")?; |
| 109 | } |
| 110 | } |
| 111 | Ok(()) |
| 112 | })) |
| 113 | .append_element_content_handler(text!("a[href]", move |text| { |
| 114 | if labels.get().is_some() { text.remove(); } |
| 115 | Ok(()) |
| 116 | })))?) |
| 117 | } |
| 118 | |
| 119 | fn strip_hop_headers(headers: &mut HeaderMap) { |
| 120 | let named: Vec<_> = headers |
| 121 | .get_all("connection") |
| 122 | .iter() |
| 123 | .filter_map(|v| v.to_str().ok()) |
| 124 | .flat_map(|v| v.split(',')) |
| 125 | .map(|v| v.trim().to_owned()) |
| 126 | .collect(); |
| 127 | for name in named { |
| 128 | headers.remove(name); |
| 129 | } |
| 130 | for name in [ |
| 131 | "connection", |
| 132 | "keep-alive", |
| 133 | "proxy-authenticate", |
| 134 | "proxy-authorization", |
| 135 | "te", |
| 136 | "trailer", |
| 137 | "transfer-encoding", |
| 138 | "upgrade", |
| 139 | ] { |
| 140 | headers.remove(name); |
| 141 | } |
| 142 | } |
| 143 | |
| 144 | pub async fn proxy(app: Arc<App>, request: Request) -> Result<Response> { |
| 145 | let (mut parts, body) = request.into_parts(); |
| 146 | let host = parts |
| 147 | .headers |
| 148 | .get("host") |
| 149 | .and_then(|v| v.to_str().ok()) |
| 150 | .unwrap_or_default(); |
| 151 | let domain = env("STUDIO_DOMAIN", "studio.test"); |
| 152 | let service = host.strip_suffix(&format!(".{domain}")).unwrap_or_default(); |
| 153 | if service != "shale" |
| 154 | && !(service |
| 155 | .strip_prefix("shale-preview-") |
| 156 | .is_some_and(|id| id.len() == 8 && id.bytes().all(|b| b.is_ascii_hexdigit()))) |
| 157 | { |
| 158 | return Err(Error::new(403, "Open Shale to continue.")); |
| 159 | } |
| 160 | let upstream: std::net::SocketAddr = parts |
| 161 | .headers |
| 162 | .get("studio-shale-upstream") |
| 163 | .and_then(|v| v.to_str().ok()) |
| 164 | .ok_or_else(|| Error::new(403, "Open Shale to continue."))? |
| 165 | .parse()?; |
| 166 | let uri: axum::http::Uri = parts |
| 167 | .headers |
| 168 | .get("studio-shale-uri") |
| 169 | .and_then(|v| v.to_str().ok()) |
| 170 | .ok_or_else(|| Error::new(403, "Open Shale to continue."))? |
| 171 | .parse()?; |
| 172 | if uri.scheme().is_some() || uri.authority().is_some() || !uri.path().starts_with('/') { |
| 173 | return Err(Error::new(403, "Open Shale to continue.")); |
| 174 | } |
| 175 | let origin = url::Url::parse(&format!("https://{host}{uri}"))?; |
| 176 | let profiles = profiles(&app.auth)?; |
| 177 | if matches!(parts.method, Method::GET | Method::HEAD) |
| 178 | && let Some(url) = account_path(uri.path()).and_then(|name| profiles.get(name)).and_then(|profile| profile.url.as_deref()) |
| 179 | { |
| 180 | return Ok(( |
| 181 | StatusCode::FOUND, |
| 182 | [ |
| 183 | ("location", url), |
| 184 | ("cache-control", "no-store"), |
| 185 | ("referrer-policy", "no-referrer"), |
| 186 | ], |
| 187 | ) |
| 188 | .into_response()); |
| 189 | } |
| 190 | let target = match &app.internal { |
| 191 | Some((base, _)) => format!("{base}services/{service}{uri}"), |
| 192 | None => format!("http://{upstream}{uri}"), |
| 193 | }; |
| 194 | if app.internal.is_some() { |
| 195 | parts.headers.remove("host"); |
| 196 | } |
| 197 | strip_hop_headers(&mut parts.headers); |
| 198 | for name in [ |
| 199 | "studio-shale-upstream", |
| 200 | "studio-shale-uri", |
| 201 | "studio-proxy-token", |
| 202 | "user-name", |
| 203 | "user-groups", |
| 204 | "user-id", |
| 205 | "if-none-match", |
| 206 | "if-modified-since", |
| 207 | "accept-encoding", |
| 208 | ] { |
| 209 | parts.headers.remove(name); |
| 210 | } |
| 211 | parts |
| 212 | .headers |
| 213 | .insert("accept-encoding", "identity".parse().unwrap()); |
| 214 | let mut upstream_response = app |
| 215 | .shale_http |
| 216 | .request(parts.method.clone(), target) |
| 217 | .headers(parts.headers) |
| 218 | .body(reqwest::Body::wrap_stream(body.into_data_stream())) |
| 219 | .send() |
| 220 | .await?; |
| 221 | let status = upstream_response.status(); |
| 222 | let mut headers = upstream_response.headers().clone(); |
| 223 | strip_hop_headers(&mut headers); |
| 224 | let html = parts.method != Method::HEAD |
| 225 | && headers |
| 226 | .get("content-type") |
| 227 | .and_then(|v| v.to_str().ok()) |
| 228 | .is_some_and(|v| { |
| 229 | v.split(';') |
| 230 | .next() |
| 231 | .unwrap_or_default() |
| 232 | .trim() |
| 233 | .eq_ignore_ascii_case("text/html") |
| 234 | }) |
| 235 | && headers |
| 236 | .get("content-encoding") |
| 237 | .is_none_or(|v| v == "identity"); |
| 238 | let mut prefix = Vec::new(); |
| 239 | if html { |
| 240 | while let Some(chunk) = upstream_response.chunk().await? { |
| 241 | prefix.extend_from_slice(&chunk); |
| 242 | if prefix.len() > 8 * 1024 * 1024 { |
| 243 | break; |
| 244 | } |
| 245 | } |
| 246 | if prefix.len() <= 8 * 1024 * 1024 |
| 247 | && let Ok(text) = std::str::from_utf8(&prefix) |
| 248 | { |
| 249 | let rewritten = rewrite(text, &origin, &profiles)?; |
| 250 | for name in [ |
| 251 | "content-length", |
| 252 | "etag", |
| 253 | "last-modified", |
| 254 | "content-md5", |
| 255 | "digest", |
| 256 | "content-digest", |
| 257 | "accept-ranges", |
| 258 | ] { |
| 259 | headers.remove(name); |
| 260 | } |
| 261 | headers.insert("cache-control", "no-store".parse().unwrap()); |
| 262 | return Ok((status, headers, rewritten).into_response()); |
| 263 | } |
| 264 | } |
| 265 | let stream = futures::stream::once(async move { Ok::<_, reqwest::Error>(Bytes::from(prefix)) }) |
| 266 | .chain(upstream_response.bytes_stream()); |
| 267 | Ok((status, headers, Body::from_stream(stream)).into_response()) |
| 268 | } |
| 269 | |
| 270 | #[cfg(test)] |
| 271 | mod tests { |
| 272 | use super::*; |
| 273 | use scraper::{Html, Selector}; |
| 274 | |
| 275 | #[test] |
| 276 | fn rewrites_only_known_local_account_links_and_escapes_provider_names() { |
| 277 | let origin = url::Url::parse("https://shale.paperclover.net/").unwrap(); |
| 278 | let profiles = HashMap::from([ |
| 279 | ( |
| 280 | "guest-github-123".into(), |
| 281 | Profile { |
| 282 | name: "<&\"clover".into(), |
| 283 | url: Some("https://github.com/clover".into()), |
| 284 | icon: include_str!("../web/sso/github.svg"), |
| 285 | picture: Some("https://avatars.githubusercontent.com/u/123?s=64".into()), |
| 286 | }, |
| 287 | ), |
| 288 | ( |
| 289 | "guest-astheno-abc".into(), |
| 290 | Profile { |
| 291 | name: "Astheno user".into(), |
| 292 | url: None, |
| 293 | icon: include_str!("astheno.svg"), |
| 294 | picture: None, |
| 295 | }, |
| 296 | ), |
| 297 | ]); |
| 298 | let html = r#"<a class="usa-nav-link" href="/~guest-github-123"><img src="avatar"><span>~guest-github-123</span></a><a href="https://shale.paperclover.net/~guest-astheno-abc/"><img src="astheno-avatar" alt="avatar">guest</a><a href="/~guest-github-unknown">unknown</a><a href="https://other.example/~guest-github-123">external</a><a href="/~guest-github-123/repo">repo</a><pre>~guest-github-123</pre>"#; |
| 299 | let rewritten = rewrite(html, &origin, &profiles).unwrap(); |
| 300 | let page = Html::parse_document(&rewritten); |
| 301 | let links: Vec<_> = page.select(&Selector::parse("body > a").unwrap()).collect(); |
| 302 | assert_eq!(links[0].attr("target"), Some("_blank")); |
| 303 | assert_eq!(links[0].attr("rel"), Some("noreferrer")); |
| 304 | assert!(links[0].attr("style").unwrap().contains("text-decoration:none")); |
| 305 | assert!(links[0].select(&Selector::parse("span[style]").unwrap()).next().unwrap().attr("style").unwrap().contains("text-decoration:underline")); |
| 306 | assert_eq!(links[1].value().name(), "a"); |
| 307 | assert_eq!(links[1].attr("target"), None); |
| 308 | assert_eq!(links[1].attr("rel"), None); |
| 309 | assert_eq!(links[1].text().collect::<String>().trim(), "Astheno user"); |
| 310 | for link in &links[..2] { |
| 311 | assert_eq!( |
| 312 | link.select(&Selector::parse("svg[aria-hidden=true]").unwrap()) |
| 313 | .count(), |
| 314 | 1 |
| 315 | ); |
| 316 | } |
| 317 | assert_eq!( |
| 318 | links[1] |
| 319 | .select(&Selector::parse("img").unwrap()) |
| 320 | .next() |
| 321 | .unwrap() |
| 322 | .attr("src"), |
| 323 | Some("astheno-avatar") |
| 324 | ); |
| 325 | assert_eq!(links[0].select(&Selector::parse("img").unwrap()).count(), 1); |
| 326 | assert_eq!( |
| 327 | links[0] |
| 328 | .select(&Selector::parse("img").unwrap()) |
| 329 | .next() |
| 330 | .unwrap() |
| 331 | .attr("src"), |
| 332 | Some("https://avatars.githubusercontent.com/u/123?s=64") |
| 333 | ); |
| 334 | assert_eq!(links[0].text().collect::<String>().trim(), "<&\"clover"); |
| 335 | assert_eq!(links[0].attr("class"), Some("usa-nav-link")); |
| 336 | assert_eq!(links[1].attr("href"), None); |
| 337 | for link in &links[2..] { |
| 338 | assert_eq!(link.attr("target"), None); |
| 339 | } |
| 340 | assert!(rewritten.contains("<pre>~guest-github-123</pre>")); |
| 341 | assert_eq!(account_path("/~guest-github-123//"), None); |
| 342 | } |
| 343 | } |