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