use std::sync::Arc;
use axum::extract::State;
use axum::http::{header, HeaderMap};
use axum::response::IntoResponse;
use axum::routing::get;
use axum::{Json, Router};
use rusqlite::Connection;
use sha2::{Digest, Sha256};
use crate::error::MatrixError;
#[cfg(any(test, feature = "local-bearer"))]
use crate::keys::CredentialKind;
use crate::live::{ClaimRateLimiter, LiveRegistry};
use crate::typing::TypingRegistry;
mod account;
mod compat;
pub mod edge_auth;
mod ephemeral;
mod federation;
pub mod fed_net;
mod keys;
mod messaging;
mod rooms;
mod push;
mod sync;
pub mod identity;
pub struct Caller {
pub user_id: i64,
pub mxid: String,
pub device_id: String,
pub claims: crate::policy::Claims,
}
pub struct Homeserver {
pub db: tesserax_store::Db,
pub readers: std::sync::OnceLock<tesserax_store::ReadPool>,
pub live: LiveRegistry,
pub typing: TypingRegistry,
pub claim_rate: ClaimRateLimiter,
pub push: crate::push::PushHub,
pub public_base_url: std::sync::OnceLock<String>,
pub federation_delegate: std::sync::OnceLock<String>,
pub federation_enabled: std::sync::OnceLock<()>,
pub federation_allow: std::sync::OnceLock<Vec<String>>,
pub anon_read: std::sync::OnceLock<()>,
pub key_fetcher: std::sync::OnceLock<Arc<dyn crate::federation::RemoteKeys>>,
pub fed_transport: std::sync::OnceLock<Arc<dyn crate::federation::FedTransport>>,
pub fed_notify: tokio::sync::Notify,
pub seam: std::sync::OnceLock<Arc<identity::Seam>>,
pub policy: std::sync::OnceLock<Arc<dyn crate::policy::PolicyHook>>,
}
impl Homeserver {
pub fn new(conn: Connection) -> Self {
let db = tesserax_store::Db::open(&tesserax_store::DbConfig::in_memory()).expect("in-memory store");
db.blocking(|c| {
*c = conn;
Ok(())
})
.expect("install connection");
Self::from_db(db)
}
pub fn from_db(db: tesserax_store::Db) -> Self {
Self {
db,
readers: std::sync::OnceLock::new(),
live: LiveRegistry::new(),
typing: TypingRegistry::new(),
claim_rate: ClaimRateLimiter::new(),
push: crate::push::PushHub::new(),
public_base_url: std::sync::OnceLock::new(),
federation_delegate: std::sync::OnceLock::new(),
federation_enabled: std::sync::OnceLock::new(),
federation_allow: std::sync::OnceLock::new(),
anon_read: std::sync::OnceLock::new(),
key_fetcher: std::sync::OnceLock::new(),
fed_transport: std::sync::OnceLock::new(),
fed_notify: tokio::sync::Notify::new(),
seam: std::sync::OnceLock::new(),
policy: std::sync::OnceLock::new(),
}
}
pub fn check_policy(&self, caller: &Caller, action: crate::policy::Action, room_kind: Option<crate::store::RoomKind>) -> Result<(), MatrixError> {
match self.policy.get() {
Some(h) => crate::policy::enforce(h.as_ref(), &caller.claims, action, room_kind),
None => Ok(()),
}
}
}
pub fn hash_token(raw: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(raw.as_bytes());
hex::encode(hasher.finalize())
}
pub fn raw_token(headers: &HeaderMap, query_token: Option<&str>) -> Option<String> {
if let Some(value) = headers.get(header::AUTHORIZATION).and_then(|v| v.to_str().ok()) {
if let Some(rest) = value.strip_prefix("Bearer ") {
let rest = rest.trim();
if !rest.is_empty() {
return Some(rest.to_string());
}
}
}
query_token.map(str::trim).filter(|s| !s.is_empty()).map(str::to_string)
}
impl Homeserver {
pub async fn conn_async<T: Send + 'static>(&self, f: impl FnOnce(&mut Connection) -> T + Send + 'static) -> T {
self.db.write(move |c| Ok(f(c))).await.expect("store writer")
}
pub fn conn_scope<T>(&self, f: impl FnOnce(&mut Connection) -> T) -> T {
self.db.write_blocking(|c| Ok(f(c))).expect("store writer")
}
}
pub async fn resolve_caller(
state: &Arc<Homeserver>,
headers: &HeaderMap,
query_token: Option<&str>,
) -> Result<Caller, MatrixError> {
if let Some((user_id, device_id, claims)) = identity::read_resolved(headers) {
if user_id == identity::ANON_USER {
return Ok(Caller { user_id, mxid: String::new(), device_id: String::new(), claims });
}
let state = Arc::clone(state);
return tokio::task::spawn_blocking(move || -> Result<Caller, MatrixError> {
state.conn_scope(|conn: &mut rusqlite::Connection| {
let mxid = crate::store::mxid_of(&conn, user_id)?.ok_or_else(MatrixError::unknown_token)?;
Ok(Caller { user_id, mxid, device_id, claims })
})
})
.await
.map_err(|_| MatrixError::internal())?;
}
#[cfg(not(any(test, feature = "local-bearer")))]
{
let _ = query_token;
return Err(MatrixError::missing_token());
}
#[cfg(any(test, feature = "local-bearer"))]
resolve_bearer(state, headers, query_token).await
}
#[cfg(any(test, feature = "local-bearer"))]
async fn resolve_bearer(state: &Arc<Homeserver>, headers: &HeaderMap, query_token: Option<&str>) -> Result<Caller, MatrixError> {
let Some(raw) = raw_token(headers, query_token) else {
return Err(MatrixError::missing_token());
};
let hash = hash_token(&raw);
let state = Arc::clone(state);
tokio::task::spawn_blocking(move || -> Result<Caller, MatrixError> {
state.conn_scope(|conn: &mut rusqlite::Connection| {
let device = crate::keys::device_for_credential(&conn, CredentialKind::Bearer, &hash)?
.ok_or_else(MatrixError::unknown_token)?;
let mxid = crate::store::mxid_of(&conn, device.user_id)?.ok_or_else(MatrixError::unknown_token)?;
Ok(Caller { user_id: device.user_id, mxid, device_id: device.device_id, claims: crate::policy::Claims::new() })
})
})
.await
.map_err(|_| MatrixError::internal())?
}
pub(crate) async fn with_conn_pub<T, F>(state: &Arc<Homeserver>, work: F) -> Result<T, MatrixError>
where
T: Send + 'static,
F: FnOnce(&mut Connection) -> Result<T, MatrixError> + Send + 'static,
{
let state = Arc::clone(state);
state.db.write(move |conn| Ok(work(conn))).await.map_err(|_| MatrixError::internal())?
}
pub async fn with_read_pub<T, F>(state: &Arc<Homeserver>, work: F) -> Result<T, MatrixError>
where
T: Send + 'static,
F: FnOnce(&Connection) -> Result<T, MatrixError> + Send + 'static,
{
match state.readers.get() {
Some(pool) => pool.read(move |conn| Ok(work(conn))).await.map_err(|_| MatrixError::internal())?,
None => state.db.read(move |conn| Ok(work(conn))).await.map_err(|_| MatrixError::internal())?,
}
}
pub fn wake_users(state: &Homeserver, ids: impl IntoIterator<Item = i64>) {
state.live.wake_many(ids.into_iter().map(|id| format!("user:{id}")));
state.fed_notify.notify_one();
}
pub fn router(state: Arc<Homeserver>) -> Router {
Router::new()
.route("/client/versions", get(versions))
.route("/client/v3/capabilities", get(capabilities))
.merge(rooms::routes())
.merge(messaging::routes())
.merge(ephemeral::routes())
.merge(account::routes())
.merge(keys::routes())
.merge(sync::routes())
.merge(push::routes())
.merge(compat::routes())
.merge(federation::routes())
.merge(identity::routes())
.fallback(unrecognized)
.layer(axum::middleware::from_fn_with_state(state.clone(), identity::assertion_layer))
.with_state(state)
}
async fn versions() -> impl IntoResponse {
Json(serde_json::json!({
"versions": ["v1.19"],
"unstable_features": {},
}))
}
async fn capabilities(
State(state): State<Arc<Homeserver>>,
headers: HeaderMap,
) -> Result<impl IntoResponse, MatrixError> {
resolve_caller(&state, &headers, None).await?;
let room_version = crate::store::MATRIX_ROOM_VERSION;
Ok(Json(serde_json::json!({
"capabilities": {
"m.room_versions": {
"default": room_version,
"available": { room_version: "stable" },
},
"m.change_password": { "enabled": false },
}
})))
}
async fn unrecognized() -> MatrixError {
MatrixError::unrecognized()
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::{to_bytes, Body};
use axum::http::{Request, StatusCode};
use tower::ServiceExt;
#[tokio::test]
async fn versions_reports_v1_19() {
let conn = Connection::open_in_memory().expect("memory");
let app = router(Arc::new(Homeserver::new(conn)));
let resp = app
.oneshot(Request::get("/client/versions").body(Body::empty()).expect("request"))
.await
.expect("response");
assert_eq!(resp.status(), StatusCode::OK);
let body = to_bytes(resp.into_body(), usize::MAX).await.expect("body");
let value: serde_json::Value = serde_json::from_slice(&body).expect("json");
assert_eq!(value, serde_json::json!({ "versions": ["v1.19"], "unstable_features": {} }));
}
#[tokio::test]
async fn unknown_path_is_unrecognized() {
let conn = Connection::open_in_memory().expect("memory");
let app = router(Arc::new(Homeserver::new(conn)));
let resp = app
.oneshot(Request::get("/client/v3/nope").body(Body::empty()).expect("request"))
.await
.expect("response");
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
let body = to_bytes(resp.into_body(), usize::MAX).await.expect("body");
let value: serde_json::Value = serde_json::from_slice(&body).expect("json");
assert_eq!(value["errcode"], "M_UNRECOGNIZED");
}
}