1use crate::*;
2use axum::body::Body;
3use lol_html::{RewriteStrSettings, element, html_content::ContentType, text};
4
5struct Profile {
6 name: String,
7 url: Option<String>,
8 icon: &'static str,
9 picture: Option<String>,
10}
11
12fn 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
57fn 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
63fn 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
119fn 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
144pub 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)]
271mod 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}