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#[derive(Debug)]
26pub(crate) struct Args<T, const ADMIN: bool = false> {
27 pub(crate) body: T,
29
30 pub(crate) cookie: CookieJar,
32
33 pub(crate) origin: Option<OwnedServerName>,
35
36 pub(crate) sender_user: Option<OwnedUserId>,
38
39 pub(crate) sender_device: Option<OwnedDeviceId>,
41
42 pub(crate) appservice_info: Option<RegistrationInfo>,
44
45 pub(crate) json_body: Option<CanonicalJsonValue>,
47}
48
49pub(crate) type ArgsAdmin<T> = Args<T, true>;
53
54#[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#[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#[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
167fn 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
195async 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
213fn 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
250fn 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
288fn 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}