1use std::sync::Arc;
10
11use axum::extract::State;
12use axum::http::{header, HeaderMap};
13use axum::response::IntoResponse;
14use axum::routing::get;
15use axum::{Json, Router};
16use rusqlite::Connection;
17use sha2::{Digest, Sha256};
18
19use crate::error::MatrixError;
20#[cfg(any(test, feature = "local-bearer"))]
21use crate::keys::CredentialKind;
22use crate::live::{ClaimRateLimiter, LiveRegistry};
23use crate::typing::TypingRegistry;
24
25mod account;
26pub(crate) mod extras;
27mod compat;
28pub mod edge_auth;
29mod ephemeral;
30mod federation;
31pub mod fed_net;
32mod keys;
33mod messaging;
34mod rooms;
35mod push;
36mod sync;
37mod sliding;
38pub mod identity;
39
40pub struct Caller {
41 pub user_id: i64,
42 pub mxid: String,
43 pub device_id: String,
44 pub claims: crate::policy::Claims,
46}
47
48pub struct Homeserver {
49 pub db: tesserax_store::Db,
52 pub readers: std::sync::OnceLock<tesserax_store::ReadPool>,
54 pub live: LiveRegistry,
55 pub typing: TypingRegistry,
56 pub claim_rate: ClaimRateLimiter,
57 pub push: crate::push::PushHub,
58 pub public_base_url: std::sync::OnceLock<String>,
60 pub federation_delegate: std::sync::OnceLock<String>,
62 pub federation_enabled: std::sync::OnceLock<()>,
64 pub federation_allow: std::sync::OnceLock<Vec<String>>,
66 pub anon_read: std::sync::OnceLock<()>,
68 pub key_fetcher: std::sync::OnceLock<Arc<dyn crate::federation::RemoteKeys>>,
70 pub fed_transport: std::sync::OnceLock<Arc<dyn crate::federation::FedTransport>>,
72 pub fed_notify: tokio::sync::Notify,
74 pub seam: std::sync::OnceLock<Arc<identity::Seam>>,
76 pub policy: std::sync::OnceLock<Arc<dyn crate::policy::PolicyHook>>,
78}
79
80impl Homeserver {
81 pub fn new(conn: Connection) -> Self {
84 let db = tesserax_store::Db::open(&tesserax_store::DbConfig::in_memory()).expect("in-memory store");
85 db.blocking(|c| {
87 *c = conn;
88 Ok(())
89 })
90 .expect("install connection");
91 Self::from_db(db)
92 }
93
94 pub fn from_db(db: tesserax_store::Db) -> Self {
96 Self {
97 db,
98 readers: std::sync::OnceLock::new(),
99 live: LiveRegistry::new(),
100 typing: TypingRegistry::new(),
101 claim_rate: ClaimRateLimiter::new(),
102 push: crate::push::PushHub::new(),
103 public_base_url: std::sync::OnceLock::new(),
104 federation_delegate: std::sync::OnceLock::new(),
105 federation_enabled: std::sync::OnceLock::new(),
106 federation_allow: std::sync::OnceLock::new(),
107 anon_read: std::sync::OnceLock::new(),
108 key_fetcher: std::sync::OnceLock::new(),
109 fed_transport: std::sync::OnceLock::new(),
110 fed_notify: tokio::sync::Notify::new(),
111 seam: std::sync::OnceLock::new(),
112 policy: std::sync::OnceLock::new(),
113 }
114 }
115
116 pub fn check_policy(&self, caller: &Caller, action: crate::policy::Action, room_kind: Option<crate::store::RoomKind>) -> Result<(), MatrixError> {
118 match self.policy.get() {
119 Some(h) => crate::policy::enforce(h.as_ref(), &caller.claims, action, room_kind),
120 None => Ok(()),
121 }
122 }
123}
124
125pub fn hash_token(raw: &str) -> String {
126 let mut hasher = Sha256::new();
127 hasher.update(raw.as_bytes());
128 hex::encode(hasher.finalize())
129}
130
131pub fn raw_token(headers: &HeaderMap, query_token: Option<&str>) -> Option<String> {
132 if let Some(value) = headers.get(header::AUTHORIZATION).and_then(|v| v.to_str().ok()) {
133 if let Some(rest) = value.strip_prefix("Bearer ") {
134 let rest = rest.trim();
135 if !rest.is_empty() {
136 return Some(rest.to_string());
137 }
138 }
139 }
140 query_token.map(str::trim).filter(|s| !s.is_empty()).map(str::to_string)
141}
142
143impl Homeserver {
144 pub async fn conn_async<T: Send + 'static>(&self, f: impl FnOnce(&mut Connection) -> T + Send + 'static) -> T {
148 self.db.write(move |c| Ok(f(c))).await.expect("store writer")
149 }
150
151 pub fn conn_scope<T>(&self, f: impl FnOnce(&mut Connection) -> T) -> T {
152 self.db.write_blocking(|c| Ok(f(c))).expect("store writer")
153 }
154}
155
156pub async fn resolve_caller(
157 state: &Arc<Homeserver>,
158 headers: &HeaderMap,
159 query_token: Option<&str>,
160) -> Result<Caller, MatrixError> {
161 if let Some((user_id, device_id, claims)) = identity::read_resolved(headers) {
162 if user_id == identity::ANON_USER {
163 return Ok(Caller { user_id, mxid: String::new(), device_id: String::new(), claims });
165 }
166 let state = Arc::clone(state);
167 return tokio::task::spawn_blocking(move || -> Result<Caller, MatrixError> {
168 state.conn_scope(|conn: &mut rusqlite::Connection| {
169 let mxid = crate::store::mxid_of(&conn, user_id)?.ok_or_else(MatrixError::unknown_token)?;
170 Ok(Caller { user_id, mxid, device_id, claims })
171 })
172 })
173 .await
174 .map_err(|_| MatrixError::internal())?;
175 }
176 #[cfg(not(any(test, feature = "local-bearer")))]
177 {
178 let _ = query_token;
179 return Err(MatrixError::missing_token());
180 }
181 #[cfg(any(test, feature = "local-bearer"))]
182 resolve_bearer(state, headers, query_token).await
183}
184
185#[cfg(any(test, feature = "local-bearer"))]
186async fn resolve_bearer(state: &Arc<Homeserver>, headers: &HeaderMap, query_token: Option<&str>) -> Result<Caller, MatrixError> {
187 let Some(raw) = raw_token(headers, query_token) else {
188 return Err(MatrixError::missing_token());
189 };
190 let hash = hash_token(&raw);
191 let state = Arc::clone(state);
192 tokio::task::spawn_blocking(move || -> Result<Caller, MatrixError> {
193 state.conn_scope(|conn: &mut rusqlite::Connection| {
194 let device = crate::keys::device_for_credential(&conn, CredentialKind::Bearer, &hash)?
195 .ok_or_else(MatrixError::unknown_token)?;
196 let mxid = crate::store::mxid_of(&conn, device.user_id)?.ok_or_else(MatrixError::unknown_token)?;
197 Ok(Caller { user_id: device.user_id, mxid, device_id: device.device_id, claims: crate::policy::Claims::new() })
198 })
199 })
200 .await
201 .map_err(|_| MatrixError::internal())?
202}
203
204pub(crate) async fn with_conn_pub<T, F>(state: &Arc<Homeserver>, work: F) -> Result<T, MatrixError>
206where
207 T: Send + 'static,
208 F: FnOnce(&mut Connection) -> Result<T, MatrixError> + Send + 'static,
209{
210 let state = Arc::clone(state);
211 state.db.write(move |conn| Ok(work(conn))).await.map_err(|_| MatrixError::internal())?
212}
213
214pub async fn with_read_pub<T, F>(state: &Arc<Homeserver>, work: F) -> Result<T, MatrixError>
217where
218 T: Send + 'static,
219 F: FnOnce(&Connection) -> Result<T, MatrixError> + Send + 'static,
220{
221 match state.readers.get() {
222 Some(pool) => pool.read(move |conn| Ok(work(conn))).await.map_err(|_| MatrixError::internal())?,
223 None => state.db.read(move |conn| Ok(work(conn))).await.map_err(|_| MatrixError::internal())?,
224 }
225}
226
227pub fn wake_users(state: &Homeserver, ids: impl IntoIterator<Item = i64>) {
228 state.live.wake_many(ids.into_iter().map(|id| format!("user:{id}")));
229 state.fed_notify.notify_one();
231}
232
233pub fn router(state: Arc<Homeserver>) -> Router {
234 Router::new()
235 .route("/client/versions", get(versions))
236 .route("/client/v3/capabilities", get(capabilities))
237 .merge(rooms::routes())
238 .merge(messaging::routes())
239 .merge(ephemeral::routes())
240 .merge(account::routes())
241 .merge(extras::routes())
242 .merge(sliding::routes())
243 .merge(keys::routes())
244 .merge(sync::routes())
245 .merge(push::routes())
246 .merge(compat::routes())
247 .merge(federation::routes())
248 .merge(identity::routes())
249 .fallback(unrecognized)
250 .layer(axum::middleware::from_fn_with_state(state.clone(), identity::assertion_layer))
251 .layer(axum::middleware::from_fn(cors))
252 .with_state(state)
253}
254
255pub async fn cors(req: axum::extract::Request, next: axum::middleware::Next) -> axum::response::Response {
257 use axum::http::{header, HeaderValue, Method, StatusCode};
258 let preflight = req.method() == Method::OPTIONS;
259 let mut resp = if preflight { axum::response::Response::builder().status(StatusCode::OK).body(axum::body::Body::empty()).unwrap() } else { next.run(req).await };
260 let h = resp.headers_mut();
261 h.insert(header::ACCESS_CONTROL_ALLOW_ORIGIN, HeaderValue::from_static("*"));
262 h.insert(header::ACCESS_CONTROL_ALLOW_METHODS, HeaderValue::from_static("GET, POST, PUT, DELETE, OPTIONS"));
263 h.insert(header::ACCESS_CONTROL_ALLOW_HEADERS, HeaderValue::from_static("X-Requested-With, Content-Type, Authorization, Date"));
264 resp
265}
266
267const SPEC_VERSIONS: [&str; 19] = ["v1.1", "v1.2", "v1.3", "v1.4", "v1.5", "v1.6", "v1.7", "v1.8", "v1.9", "v1.10", "v1.11", "v1.12", "v1.13", "v1.14", "v1.15", "v1.16", "v1.17", "v1.18", "v1.19"];
269
270async fn versions() -> impl IntoResponse {
271 Json(serde_json::json!({
272 "versions": SPEC_VERSIONS,
273 "unstable_features": { "org.matrix.simplified_msc3575": true, "org.matrix.msc4186": true },
274 }))
275}
276
277async fn capabilities(
278 State(state): State<Arc<Homeserver>>,
279 headers: HeaderMap,
280) -> Result<impl IntoResponse, MatrixError> {
281 resolve_caller(&state, &headers, None).await?;
282 let room_version = crate::store::MATRIX_ROOM_VERSION;
283 Ok(Json(serde_json::json!({
284 "capabilities": {
285 "m.room_versions": {
286 "default": room_version,
287 "available": { room_version: "stable" },
288 },
289 "m.change_password": { "enabled": false },
290 }
291 })))
292}
293
294async fn unrecognized() -> MatrixError {
295 MatrixError::unrecognized()
296}
297
298#[cfg(test)]
299mod tests {
300 use super::*;
301 use axum::body::{to_bytes, Body};
302 use axum::http::{Request, StatusCode};
303 use tower::ServiceExt;
304
305 #[tokio::test]
306 async fn versions_reports_v1_19() {
307 let conn = Connection::open_in_memory().expect("memory");
308 let app = router(Arc::new(Homeserver::new(conn)));
309 let resp = app
310 .oneshot(Request::get("/client/versions").body(Body::empty()).expect("request"))
311 .await
312 .expect("response");
313 assert_eq!(resp.status(), StatusCode::OK);
314 let body = to_bytes(resp.into_body(), usize::MAX).await.expect("body");
315 let value: serde_json::Value = serde_json::from_slice(&body).expect("json");
316 assert_eq!(value["versions"].as_array().unwrap().last().unwrap(), "v1.19");
317 assert!(value["versions"].as_array().unwrap().iter().any(|v| v == "v1.1"));
318 }
319
320 #[tokio::test]
321 async fn unknown_path_is_unrecognized() {
322 let conn = Connection::open_in_memory().expect("memory");
323 let app = router(Arc::new(Homeserver::new(conn)));
324 let resp = app
325 .oneshot(Request::get("/client/v3/nope").body(Body::empty()).expect("request"))
326 .await
327 .expect("response");
328 assert_eq!(resp.status(), StatusCode::NOT_FOUND);
329 let body = to_bytes(resp.into_body(), usize::MAX).await.expect("body");
330 let value: serde_json::Value = serde_json::from_slice(&body).expect("json");
331 assert_eq!(value["errcode"], "M_UNRECOGNIZED");
332 }
333}
334