diff --git a/src/bot/implementation.rs b/src/bot/implementation.rs index 543fb1e..0a75694 100644 --- a/src/bot/implementation.rs +++ b/src/bot/implementation.rs @@ -6,7 +6,7 @@ use mxlink::matrix_sdk::Room; use mxlink::matrix_sdk::media::{MediaFormat, MediaRequestParameters}; use mxlink::matrix_sdk::ruma::api::client::profile::{AvatarUrl, DisplayName}; use mxlink::matrix_sdk::ruma::{ - MilliSecondsSinceUnixEpoch, OwnedDeviceId, OwnedUserId, events::room::MediaSource, + MilliSecondsSinceUnixEpoch, OwnedUserId, events::room::MediaSource, }; use mxlink::{ @@ -25,7 +25,7 @@ use crate::agent::Manager as AgentManager; use crate::entity::catch_up_marker::{ CatchUpMarker, CatchUpMarkerManager, DelayedCatchUpMarkerManager, }; -use crate::entity::cfg::{Avatar, Config}; +use crate::entity::cfg::{Avatar, Config, ConfigUserAuth}; use crate::entity::globalconfig::{GlobalConfig, GlobalConfigurationManager}; use crate::entity::roomconfig::{RoomConfig, RoomConfigurationManager}; @@ -395,44 +395,21 @@ async fn create_matrix_link(config: &Config) -> anyhow::Result { let session_encryption_key = config.persistence.session_encryption_key()?; let db_dir_path: std::path::PathBuf = config.persistence.db_dir_path()?; - let login_creds = if let Some(access_token) = config - .user - .access_token - .as_deref() - .filter(|token| !token.is_empty()) - { - let server_name = &config.homeserver.server_name; - let localpart = &config.user.mxid_localpart; - let user_id = OwnedUserId::try_from(format!("@{localpart}:{server_name}")) - .map_err(|e| anyhow::anyhow!("Invalid user ID: {e}"))?; - let device_id = OwnedDeviceId::from( - config - .user - .device_id - .as_deref() - .filter(|device_id| !device_id.is_empty()) - .ok_or_else(|| { - anyhow::anyhow!("user.device_id must be set for access token authentication") - })?, - ); - LoginCredentials::AccessToken { + let user_auth = config.user.auth_config(&config.homeserver.server_name)?; + + let login_creds = match user_auth { + ConfigUserAuth::UserPassword { username, password } => { + LoginCredentials::UserPassword(username, password) + } + ConfigUserAuth::AccessToken { user_id, device_id, - access_token: access_token.to_owned(), - } - } else { - LoginCredentials::UserPassword( - config.user.mxid_localpart.to_owned(), - config - .user - .password - .as_deref() - .filter(|password| !password.is_empty()) - .ok_or_else(|| { - anyhow::anyhow!("user.password must be set for password authentication") - })? - .to_owned(), - ) + access_token, + } => LoginCredentials::AccessToken { + user_id, + device_id, + access_token, + }, }; let login_encryption = LoginEncryption::new( diff --git a/src/bot/load_config.rs b/src/bot/load_config.rs index 740138e..eed46b2 100644 --- a/src/bot/load_config.rs +++ b/src/bot/load_config.rs @@ -29,9 +29,15 @@ pub fn load() -> anyhow::Result { cfg_env::BAIBOT_HOMESERVER_SERVER_NAME => config.homeserver.server_name = value, cfg_env::BAIBOT_HOMESERVER_URL => config.homeserver.url = value, cfg_env::BAIBOT_USER_MXID_LOCALPART => config.user.mxid_localpart = value, - cfg_env::BAIBOT_USER_PASSWORD => config.user.password = Some(value), - cfg_env::BAIBOT_USER_ACCESS_TOKEN => config.user.access_token = Some(value), - cfg_env::BAIBOT_USER_DEVICE_ID => config.user.device_id = Some(value), + cfg_env::BAIBOT_USER_PASSWORD => { + config.user.password = optional_non_empty(value); + } + cfg_env::BAIBOT_USER_ACCESS_TOKEN => { + config.user.access_token = optional_non_empty(value); + } + cfg_env::BAIBOT_USER_DEVICE_ID => { + config.user.device_id = optional_non_empty(value); + } cfg_env::BAIBOT_USER_ENCRYPTION_RECOVERY_PASSPHRASE => { config.user.encryption.recovery_passphrase = Some(value); } @@ -122,3 +128,7 @@ pub fn load() -> anyhow::Result { Ok(config) } + +fn optional_non_empty(value: String) -> Option { + if value.is_empty() { None } else { Some(value) } +} diff --git a/src/entity/cfg/config.rs b/src/entity/cfg/config.rs index c17d344..b74643e 100644 --- a/src/entity/cfg/config.rs +++ b/src/entity/cfg/config.rs @@ -1,6 +1,7 @@ use std::path::PathBuf; use mxlink::helpers::encryption::EncryptionKey; +use mxlink::matrix_sdk::ruma::{OwnedDeviceId, OwnedUserId}; use serde::{Deserialize, Deserializer, Serialize}; use crate::{ @@ -38,7 +39,7 @@ pub struct Config { impl Config { pub fn validate(&self) -> anyhow::Result<()> { self.homeserver.validate()?; - self.user.validate()?; + self.user.validate(&self.homeserver.server_name)?; self.persistence.validate()?; self.room.validate()?; self.access.validate()?; @@ -57,6 +58,19 @@ impl Config { } } +#[derive(Debug)] +pub enum ConfigUserAuth { + UserPassword { + username: String, + password: String, + }, + AccessToken { + user_id: OwnedUserId, + device_id: OwnedDeviceId, + access_token: String, + }, +} + #[derive(Debug, Serialize, Deserialize)] pub struct ConfigHomeserver { pub server_name: String, @@ -148,7 +162,7 @@ pub struct ConfigUser { } impl ConfigUser { - pub fn validate(&self) -> anyhow::Result<()> { + pub fn validate(&self, homeserver_server_name: &str) -> anyhow::Result<()> { if self.mxid_localpart.is_empty() { return Err(anyhow::anyhow!( "The user.mxid_localpart ({}) configuration must be set", @@ -156,7 +170,7 @@ impl ConfigUser { )); } - self.validate_auth()?; + self.auth_config(homeserver_server_name)?; if self.name.is_empty() { return Err(anyhow::anyhow!( @@ -170,27 +184,55 @@ impl ConfigUser { Ok(()) } - fn validate_auth(&self) -> anyhow::Result<()> { - let has_password = self.password.as_deref().is_some_and(|p| !p.is_empty()); - let has_access_token = self.access_token.as_deref().is_some_and(|t| !t.is_empty()); - let has_device_id = self.device_id.as_deref().is_some_and(|d| !d.is_empty()); + pub fn auth_config(&self, homeserver_server_name: &str) -> anyhow::Result { + let password = self.password.as_deref().filter(|value| !value.is_empty()); + let access_token = self + .access_token + .as_deref() + .filter(|value| !value.is_empty()); - if !has_password && !has_access_token { - return Err(anyhow::anyhow!( - "Either user.password ({}) or user.access_token ({}) must be set", + match (password, access_token) { + (Some(_), Some(_)) => Err(anyhow::anyhow!( + "Set exactly one authentication method: either user.password ({}) OR user.access_token ({}) + user.device_id ({})", super::env::BAIBOT_USER_PASSWORD, - super::env::BAIBOT_USER_ACCESS_TOKEN - )); - } - - if has_access_token && !has_device_id { - return Err(anyhow::anyhow!( - "user.device_id ({}) must be set when using access token authentication", + super::env::BAIBOT_USER_ACCESS_TOKEN, super::env::BAIBOT_USER_DEVICE_ID - )); - } + )), + (None, None) => Err(anyhow::anyhow!( + "Set one authentication method: either user.password ({}) OR user.access_token ({}) + user.device_id ({})", + super::env::BAIBOT_USER_PASSWORD, + super::env::BAIBOT_USER_ACCESS_TOKEN, + super::env::BAIBOT_USER_DEVICE_ID + )), + (Some(password), None) => Ok(ConfigUserAuth::UserPassword { + username: self.mxid_localpart.to_owned(), + password: password.to_owned(), + }), + (None, Some(access_token)) => { + let device_id = self + .device_id + .as_deref() + .filter(|value| !value.is_empty()) + .ok_or_else(|| { + anyhow::anyhow!( + "user.device_id ({}) must be set when using access token authentication", + super::env::BAIBOT_USER_DEVICE_ID + ) + })?; - Ok(()) + let user_id = OwnedUserId::try_from(format!( + "@{}:{}", + self.mxid_localpart, homeserver_server_name + )) + .map_err(|e| anyhow::anyhow!("Invalid user ID: {e}"))?; + + Ok(ConfigUserAuth::AccessToken { + user_id, + device_id: OwnedDeviceId::from(device_id), + access_token: access_token.to_owned(), + }) + } + } } } diff --git a/src/entity/cfg/mod.rs b/src/entity/cfg/mod.rs index b07a737..5902b9a 100644 --- a/src/entity/cfg/mod.rs +++ b/src/entity/cfg/mod.rs @@ -2,4 +2,4 @@ mod config; pub mod defaults; pub mod env; -pub use config::{Avatar, Config}; +pub use config::{Avatar, Config, ConfigUserAuth};