Skip to main content

tuwunel_api/router/
args.rs

1use std::{any::TypeId, fmt::Debug, ops::Deref};
2
3use axum::{body::Body, extract::FromRequest};
4use axum_extra::extract::cookie::CookieJar;
5use bytes::{BufMut, Bytes, BytesMut};
6use http::{Method, Request as HttpRequest};
7use ruma::{
8	CanonicalJsonObject, CanonicalJsonValue, DeviceId, OwnedDeviceId, OwnedServerName,
9	OwnedUserId, ServerName, UserId, api::IncomingRequest,
10};
11use serde_json::{Value as JsonValue, from_slice};
12use tuwunel_core::{Error, Result, err, implement, utils::string::EMPTY};
13use tuwunel_service::{Services, appservice::RegistrationInfo};
14
15use super::{
16	auth::{Auth, AuthDispatch, auth},
17	request::{Request, from as request_from},
18};
19use crate::{State, client::admin::require_admin};
20
21/// Extracts a typed Ruma request and its authentication context.
22///
23/// Administrator-only routes defer JSON errors until authorization succeeds.
24/// Ordinary routes reject malformed JSON before authentication.
25#[derive(Debug)]
26pub(crate) struct Args<T, const ADMIN: bool = false> {
27	/// Request struct body
28	pub(crate) body: T,
29
30	/// Cookies received from the useragent.
31	pub(crate) cookie: CookieJar,
32
33	/// Authenticated X-Matrix origin, absent for non-federation requests.
34	pub(crate) origin: Option<OwnedServerName>,
35
36	/// Authenticated local user, absent when no local user is identified.
37	pub(crate) sender_user: Option<OwnedUserId>,
38
39	/// Authenticated local device, absent for device-less authentication.
40	pub(crate) sender_device: Option<OwnedDeviceId>,
41
42	/// Authenticated appservice registration, absent for other callers.
43	pub(crate) appservice_info: Option<RegistrationInfo>,
44
45	/// Parsed canonical JSON, absent for raw or noncanonical request bodies.
46	pub(crate) json_body: Option<CanonicalJsonValue>,
47}
48
49/// Requires administrator authorization before returning request body errors.
50///
51/// Routes opt in through this alias; the default extractor retains its ordering.
52pub(crate) type ArgsAdmin<T> = Args<T, true>;
53
54/// Returns the user authenticated for a route requiring a user identity.
55///
56/// Panics if the endpoint's authentication scheme did not identify a user.
57#[implement(
58	Args,
59	generics = "<T, const ADMIN: bool>",
60	params = "<T, ADMIN>"
61)]
62#[inline]
63pub(crate) fn sender_user(&self) -> &UserId {
64	self.sender_user
65		.as_deref()
66		.expect("user must be authenticated for this handler")
67}
68
69/// Returns the server authenticated for a federation route.
70///
71/// Panics if the endpoint's authentication scheme did not identify a server.
72#[implement(
73	Args,
74	generics = "<T, const ADMIN: bool>",
75	params = "<T, ADMIN>"
76)]
77#[inline]
78pub(crate) fn origin(&self) -> &ServerName {
79	self.origin
80		.as_deref()
81		.expect("server must be authenticated for this handler")
82}
83
84/// Returns the authenticated device or rejects a device-less request.
85///
86/// User authentication alone does not guarantee a device identity.
87#[implement(
88	Args,
89	generics = "<T, const ADMIN: bool>",
90	params = "<T, ADMIN>"
91)]
92#[inline]
93pub(crate) fn sender_device(&self) -> Result<&DeviceId> {
94	self.sender_device
95		.as_deref()
96		.ok_or(err!(Request(Forbidden("user must be authenticated and device identified"))))
97}
98
99impl<T, const ADMIN: bool> Deref for Args<T, ADMIN>
100where
101	T: Sync,
102{
103	type Target = T;
104
105	fn deref(&self) -> &Self::Target { &self.body }
106}
107
108impl<T, const ADMIN: bool> FromRequest<State, Body> for Args<T, ADMIN>
109where
110	T: IncomingRequest + Debug + Send + Sync + 'static,
111	T::Authentication: AuthDispatch,
112{
113	type Rejection = Error;
114
115	#[tracing::instrument(name = "ar", level = "debug", skip_all, err(level = "debug"))]
116	async fn from_request(
117		request: HttpRequest<Body>,
118		services: &State,
119	) -> Result<Self, Self::Rejection> {
120		let request = request_from(services, request).await?;
121		let json_body = match ADMIN {
122			| true => parse_json(&request),
123			| false => Ok(parse_json(&request)?),
124		};
125
126		let json = json_body.as_ref().ok().and_then(Option::as_ref);
127		let (request, auth) = authenticate::<T>(request, services, json).await?;
128
129		_ = request
130			.parts
131			.extensions
132			.get::<tracing::Span>()
133			.inspect(|span| record_auth_context(span, &auth));
134
135		if ADMIN {
136			let sender = auth.sender_user.as_deref().ok_or_else(|| {
137				err!(Request(Forbidden("Only server administrators can use this endpoint")))
138			})?;
139
140			require_admin(services, sender).await?;
141		}
142
143		make_args(services, request, json_body?, auth)
144	}
145}
146
147fn record_auth_context(span: &tracing::Span, auth: &Auth) {
148	_ = auth
149		.sender_user
150		.as_deref()
151		.inspect(|sender_user| {
152			span.record("user_id", sender_user.as_str());
153		});
154
155	_ = auth
156		.sender_device
157		.as_deref()
158		.inspect(|sender_device| {
159			span.record("device_id", sender_device.as_str());
160		});
161
162	_ = auth.origin.as_deref().inspect(|origin| {
163		span.record("origin", origin.as_str());
164	});
165}
166
167/// Parses canonical JSON while retaining ordinary JSON for typed deserialization.
168///
169/// Empty POST and DELETE bodies become objects for UIA. Other methods and media
170/// uploads preserve their existing body handling.
171fn parse_json(request: &Request) -> Result<Option<CanonicalJsonValue>> {
172	let json_body = from_slice(&request.body).ok();
173	let json_endpoint = matches!(
174		request.parts.method,
175		Method::POST | Method::PUT | Method::DELETE | Method::PATCH
176	) && !request.parts.uri.path().contains("/media/");
177
178	if json_body.is_some() || !json_endpoint {
179		return Ok(json_body);
180	}
181
182	let empty = request.body.iter().all(u8::is_ascii_whitespace);
183
184	if !empty {
185		from_slice::<JsonValue>(&request.body)
186			.map_err(|_| err!(Request(NotJson("Request body is not valid JSON."))))?;
187	}
188
189	let empty_object = (empty && matches!(request.parts.method, Method::POST | Method::DELETE))
190		.then(|| CanonicalJsonValue::Object(CanonicalJsonObject::new()));
191
192	Ok(empty_object)
193}
194
195/// Authenticates a request while retaining the headers consumed by extraction.
196///
197/// The parsed canonical body remains available to federation authentication.
198async fn authenticate<T>(
199	mut request: Request,
200	services: &State,
201	json_body: Option<&CanonicalJsonValue>,
202) -> Result<(Request, Auth)>
203where
204	T: IncomingRequest + Debug + Send + Sync + 'static,
205	T::Authentication: AuthDispatch,
206{
207	let auth =
208		auth::<T::Authentication>(services, &mut request, json_body, TypeId::of::<T>()).await?;
209
210	Ok((request, auth))
211}
212
213/// Builds typed arguments after authentication and any UIA body merge.
214///
215/// Transfers the original HTTP parts without cloning headers or the URI.
216fn make_args<T, const ADMIN: bool>(
217	services: &Services,
218	request: Request,
219	json_body: Option<CanonicalJsonValue>,
220	auth: Auth,
221) -> Result<Args<T, ADMIN>>
222where
223	T: IncomingRequest,
224{
225	let json_body = json_body.map(|json| match json {
226		| CanonicalJsonValue::Object(json) => restore_body(services, json, &auth).into(),
227		| json => json,
228	});
229
230	let body = json_body
231		.as_ref()
232		.filter(|json| json.is_object())
233		.map_or(request.body, serialize_body);
234
235	let http_request = HttpRequest::from_parts(request.parts, body);
236	let body = T::try_from_http_request(http_request, &request.path)
237		.map_err(|e| err!(Request(BadJson(debug_warn!("{e}")))))?;
238
239	Ok(Args {
240		body,
241		cookie: request.cookie,
242		origin: auth.origin,
243		sender_user: auth.sender_user,
244		sender_device: auth.sender_device,
245		appservice_info: auth.appservice_info,
246		json_body,
247	})
248}
249
250/// Restores omitted fields from a UIA session before typed body parsing.
251///
252/// Current request fields take precedence over the saved initial request.
253fn restore_body(
254	services: &Services,
255	json_body: CanonicalJsonObject,
256	auth: &Auth,
257) -> CanonicalJsonObject {
258	let uiaa_request = json_body
259		.get("auth")
260		.and_then(CanonicalJsonValue::as_object)
261		.and_then(|auth| auth.get("session"))
262		.and_then(CanonicalJsonValue::as_str)
263		.and_then(|session| {
264			let user_id = auth.sender_user.clone().unwrap_or_else(|| {
265				UserId::parse_with_server_name(EMPTY, services.globals.server_name())
266					.expect("valid user_id")
267			});
268
269			services
270				.uiaa
271				.get_uiaa_request(&user_id, auth.sender_device.as_deref(), session)
272		});
273
274	uiaa_request
275		.and_then(|json| match json {
276			| CanonicalJsonValue::Object(json) => Some(json),
277			| _ => None,
278		})
279		.into_iter()
280		.flatten()
281		.fold(json_body, |mut json, (key, value)| {
282			json.entry(key).or_insert(value);
283
284			json
285		})
286}
287
288/// Serializes a canonical body directly into the HTTP byte buffer.
289///
290/// Canonical JSON values are always serializable, including restored UIA bodies.
291fn serialize_body(json_body: &CanonicalJsonValue) -> Bytes {
292	let mut buf = BytesMut::new().writer();
293
294	serde_json::to_writer(&mut buf, json_body).expect("value serialization can't fail");
295
296	buf.into_inner().freeze()
297}