pub mod cli;
mod cmd;
#[cfg(feature = "wake-grok")]
mod grok_listen;
mod ipc;
mod machine;
#[cfg(feature = "local-bus")]
mod local_bus;
#[cfg(not(feature = "local-bus"))]
#[path = "local_bus_off.rs"]
mod local_bus;
mod nick;
mod webhook;
#[cfg(feature = "wake-grok")]
mod node;
pub mod provider;
mod bus_backend;
pub mod resolve;
mod send;
pub mod store_key;
use std::collections::HashSet;
use std::path::{Component, Path, PathBuf};
use std::sync::Arc;
use std::time::Duration;
use mail4agent_messenger::store::sealed::SealedRecordCodec;
use mail4agent_messenger::wire::Membership;
use mail4agent_messenger::{
CoreConfig, CoreSecrets, ItemContent, Jitter, MessengerCore,
MessengerError, OutgoingRequest, OutgoingRequestKind, RecordKey, SealedRecord, SendState,
StoreError,
};
use sha2::{Digest, Sha256};
use m4a_agent::engine::Driver;
use m4a_agent::{AgentError, Backend};
use zeroize::Zeroizing;
pub use machine::{
ensure_agent_webhook_routines, ensure_agent_webhook_routines_from_env, load_agents_dir,
load_session_records, HostSession, MachineClient, RoutineReport, TickReport, WakeOptions,
WakeStatus, AGENTS_DIR_ENV, AGENT_RESCAN_SECS_ENV, DEFAULT_AGENTS_DIR,
PROFILE_NOTE_ENV, SESSIONS_DIR_ENV, SESSION_IDS_ENV, SKIP_NICKS_ENV, STORE_LOCK_FILE,
WAKE_KEYCHAIN_FILE,
};
pub use mail4agent_messenger::{
CreateRoomKind, DeviceId, MessageKind, MessengerCommand, OutgoingMessage, RoomId, RoomKind,
UserId,
};
pub use cmd::{CmdReply, CmdRequest};
pub use mail4agent_messenger::EventId;
#[cfg(feature = "wake-grok")]
pub use grok_listen::{hear, GrokListener, Heard, ListenReport};
pub use nick::{nick_from_display_name, routine_folder_id};
#[cfg(feature = "wake-grok")]
pub use node::{NodeClient, NodeTickReport, NODE_DEFAULT_SOCK_NAME};
pub use provider::{
plan_chain, HostEnv, HookFlavor, ResumeSpawnAdapter, SessionRecord, WakeChain, WebVendor,
adapter_for, wake_prompt, AdapterConfig, ClaudeChannelAdapter, ClaudeRoutineFireAdapter,
CodexAppServerAdapter, CodexCloudAdapter, CodexEndpoint, CursorAgentAdapter, GrokLeaderAdapter,
KimiServerAdapter, NoInboundAdapter, ProviderKind, ProviderSession, RoutineWebhookAdapter,
SessionKind, Surface, WakeAdapter, WakeError, WakeLetter, WakeOutcome, INBOX_DIR_ENV,
PROVIDER_ENV,
};
pub use m4a_agent::engine::PushedRoomEvent;
pub use send::{
load_env_file, load_env_file_named, send_cmd_via_socket, send_sock_path, send_sock_path_named,
send_via_socket,
SendReply, SendRequest, DEFAULT_SOCK_NAME, ENV_FILE_ENV, MAX_SEND_BYTES, SEND_SOCK_ENV,
};
pub fn store_seal_key(session_id: &str) -> [u8; 32] {
let digest = Sha256::digest(session_id.as_bytes());
let mut key = [0u8; 32];
key.copy_from_slice(&digest);
key
}
pub const STORE_ROOT_ENV: &str = "M4A_STORE_ROOT";
pub const HOMESERVER_URL_ENV: &str = "M4A_HOMESERVER_URL";
pub const CONFIG_ENV: &str = "M4A_CONFIG";
pub const BOT_NAME_ENV: &str = "M4A_BOT_NAME";
pub const SESSION_ID_ENV: &str = "M4A_SESSION_ID";
pub const PRODUCT_URL_ENV: &str = "M4A_PRODUCT_URL";
pub const PRODUCT_INVITE_ENV: &str = "M4A_PRODUCT_INVITE";
pub const TIER_ENV: &str = "M4A_TIER";
#[derive(Clone)]
#[cfg_attr(not(any(feature = "tier-server", feature = "tier-matrix")), allow(dead_code))]
struct IdentityAuth {
tier: m4a_agent::BackendKind,
invite: Option<Zeroizing<String>>,
}
pub fn session_store_dir(root: &Path, session_id: &str) -> PathBuf {
root.join(hex_encode(&store_seal_key(session_id)))
}
pub fn store_root() -> Result<PathBuf, ShellError> {
if let Some(root) = nonempty_var(STORE_ROOT_ENV) {
return Ok(PathBuf::from(root));
}
#[cfg(test)]
{
return Ok(std::env::temp_dir().join(format!(
"mail4agent-messenger-shell-root-{}",
std::process::id()
)));
}
#[cfg(not(test))]
{
Err(ShellError::StoreRoot)
}
}
fn hex_encode(bytes: &[u8]) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut out = String::with_capacity(bytes.len() * 2);
for byte in bytes {
out.push(HEX[(byte >> 4) as usize] as char);
out.push(HEX[(byte & 0x0f) as usize] as char);
}
out
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct WakeAttempt {
pub event_id: String,
pub status: Option<u16>,
}
pub struct RoutineWake {
pub url: String,
}
pub const ROUTINE_URL_ENV: &str = "M4A_ROUTINE_URL";
pub const ROUTINE_BEARER_ENV: &str = "M4A_ROUTINE_BEARER";
pub const LEADER_SOCK_ENV: &str = "M4A_LEADER_SOCK";
pub const LEADER_CWD_ENV: &str = "M4A_LEADER_CWD";
pub struct SessionWake {
pub routine_url: Option<String>,
pub routine_bearer: Option<String>,
pub leader_sock: Option<PathBuf>,
pub leader_cwd: Option<String>,
}
impl Default for SessionWake {
fn default() -> Self {
Self {
routine_url: None,
routine_bearer: None,
leader_sock: None,
leader_cwd: None,
}
}
}
impl SessionWake {
pub fn web_host() -> Self {
Self::web_from_lookup(|key| nonempty_var(key))
}
pub fn web_from_lookup(mut get: impl FnMut(&str) -> Option<String>) -> Self {
let routine_url = get(ROUTINE_URL_ENV).filter(|value| !value.is_empty());
let routine_bearer = if routine_url.is_some() {
get(ROUTINE_BEARER_ENV).filter(|value| !value.is_empty())
} else {
None
};
Self {
routine_url,
routine_bearer,
leader_sock: None,
leader_cwd: None,
}
}
pub fn node_cli() -> Result<Self, ShellError> {
Self::node_from_lookup(|key| nonempty_var(key))
}
pub fn node_from_lookup(
mut get: impl FnMut(&str) -> Option<String>,
) -> Result<Self, ShellError> {
let routine_url = get(ROUTINE_URL_ENV).filter(|value| !value.is_empty());
let routine_bearer = get(ROUTINE_BEARER_ENV).filter(|value| !value.is_empty());
if routine_url.is_some() || routine_bearer.is_some() {
return Err(ShellError::NodeRoutine);
}
Ok(Self {
routine_url: None,
routine_bearer: None,
leader_sock: get(LEADER_SOCK_ENV)
.filter(|value| !value.is_empty())
.map(PathBuf::from),
leader_cwd: get(LEADER_CWD_ENV).filter(|value| !value.is_empty()),
})
}
}
pub(crate) fn nonempty_var(name: &str) -> Option<String> {
std::env::var(name).ok().filter(|value| !value.is_empty())
}
pub struct SessionConfig {
homeserver_url: String,
session_id: String,
store_root: PathBuf,
identity: IdentityAuth,
}
impl SessionConfig {
pub fn new_identity(
product_url: impl Into<String>,
tier: m4a_agent::BackendKind,
session_id: impl Into<String>,
store_root: impl Into<PathBuf>,
invite: Option<String>,
) -> Result<Self, ShellError> {
let homeserver_url = product_url.into();
parse_base_url(&homeserver_url)?;
let session_id = session_id.into();
validate_session_id(&session_id)?;
let store_root = store_root.into();
let invite = invite.filter(|i| !i.is_empty()).or_else(|| read_invite_file(&store_root, &session_id));
Ok(Self {
homeserver_url,
session_id,
store_root,
identity: IdentityAuth { tier, invite: invite.map(Zeroizing::new) },
})
}
pub fn for_session(url: &str, session_id: &str, store_root: &Path) -> Result<Self, ShellError> {
let tier = match nonempty_var(TIER_ENV).as_deref() {
None | Some("server") => m4a_agent::BackendKind::Server,
Some("matrix") => m4a_agent::BackendKind::Matrix,
Some(_) => return Err(ShellError::Register(format!("{TIER_ENV} is server or matrix"))),
};
Self::new_identity(url, tier, session_id, store_root, None)
}
pub fn from_env() -> Result<Self, ShellError> {
let toml_text = load_homeserver_toml()?;
Self::from_lookup(
|key| std::env::var(key).ok().filter(|value| !value.is_empty()),
toml_text.as_deref(),
)
}
pub(crate) fn from_lookup(
mut get: impl FnMut(&str) -> Option<String>,
toml_text: Option<&str>,
) -> Result<Self, ShellError> {
let from_toml = match toml_text {
Some(text) => homeserver_url_from_toml(text)?,
None => None,
};
let url = get(PRODUCT_URL_ENV)
.or_else(|| get(HOMESERVER_URL_ENV))
.or(from_toml)
.ok_or(ShellError::HomeserverUrl)?;
let tier = match get(TIER_ENV).as_deref() {
None | Some("server") => m4a_agent::BackendKind::Server,
Some("matrix") => m4a_agent::BackendKind::Matrix,
Some(_) => return Err(ShellError::Register(format!("{TIER_ENV} is server or matrix"))),
};
let session_id = get(SESSION_ID_ENV).ok_or(ShellError::EmptySession)?;
let store_root = get(STORE_ROOT_ENV).ok_or(ShellError::StoreRoot)?;
Self::new_identity(url, tier, session_id, store_root, get(PRODUCT_INVITE_ENV))
}
pub fn homeserver_url(&self) -> &str {
&self.homeserver_url
}
pub fn store_dir(&self) -> PathBuf {
session_store_dir(&self.store_root, &self.session_id)
}
pub fn session_id(&self) -> &str {
&self.session_id
}
}
impl std::fmt::Debug for SessionConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SessionConfig")
.field("homeserver_url", &self.homeserver_url)
.field("session_id", &self.session_id)
.field("store_root", &self.store_root)
.field("tier", &self.identity.tier)
.finish()
}
}
fn invite_path(store_root: &Path, session_id: &str) -> PathBuf {
store_root.join("invites").join(hex_encode(&store_seal_key(session_id)))
}
fn read_invite_file(store_root: &Path, session_id: &str) -> Option<String> {
let text = std::fs::read_to_string(invite_path(store_root, session_id)).ok()?;
let code = text.trim().to_string();
(!code.is_empty()).then_some(code)
}
fn validate_session_id(session_id: &str) -> Result<(), ShellError> {
if session_id.is_empty()
|| session_id.len() > 128
|| session_id.starts_with("legacy-user-")
|| session_id
.chars()
.any(|ch| ch.is_whitespace() || ch.is_control())
{
return Err(ShellError::EmptySession);
}
Ok(())
}
fn load_homeserver_toml() -> Result<Option<String>, ShellError> {
let path = if let Some(configured) = nonempty_var(CONFIG_ENV) {
PathBuf::from(configured)
} else {
let cwd = PathBuf::from("mail4agent.toml");
if !cwd.exists() {
return Ok(None);
}
cwd
};
if !path.is_file() {
return Err(ShellError::Config("config file is missing".to_string()));
}
std::fs::read_to_string(&path)
.map(Some)
.map_err(|err| ShellError::Config(clip_public(err.to_string())))
}
fn homeserver_url_from_toml(text: &str) -> Result<Option<String>, ShellError> {
#[derive(serde::Deserialize)]
struct File {
#[serde(default)]
homeserver_url: Option<String>,
}
let file: File =
toml::from_str(text).map_err(|err| ShellError::Config(clip_public(err.to_string())))?;
Ok(file
.homeserver_url
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()))
}
pub struct FoundSession {
pub nick: String,
pub user_id: String,
}
#[derive(Default)]
pub struct DecryptedWake<'a> {
pub body: &'a str,
pub from: &'a str,
pub nick: Option<&'a str>,
pub event_id: &'a str,
pub room: Option<&'a str>,
pub from_nick: Option<&'a str>,
pub to: Option<&'a str>,
pub reply: Option<&'a str>,
}
pub const SEND_COMMAND: &str = "m4a-send";
pub fn reply_hint(to: &str, from_nick: &str) -> String {
format!("{SEND_COMMAND} --as {to} --to {from_nick} '<your reply>'")
}
pub fn mxid_localpart(mxid: &str) -> &str {
let rest = mxid.strip_prefix('@').unwrap_or(mxid);
rest.split_once(':').map(|(local, _)| local).unwrap_or(rest)
}
pub fn post_decrypted(url: &str, wake: &DecryptedWake<'_>) -> Result<(), ShellError> {
post_decrypted_with_bearer(url, wake, None)
}
pub fn post_decrypted_with_bearer(
url: &str,
wake: &DecryptedWake<'_>,
bearer: Option<&str>,
) -> Result<(), ShellError> {
let bytes = routine_json(wake)?;
post_routine_bytes(url, bytes, bearer)
}
pub fn post_routine_json(
url: &str,
body: &serde_json::Value,
bearer: Option<&str>,
) -> Result<(), ShellError> {
let bytes = serde_json::to_vec(body)
.map_err(|err| ShellError::RoutineTransport(clip_public(err.to_string())))?;
post_routine_bytes(url, bytes, bearer)
}
fn post_routine_bytes(url: &str, bytes: Vec<u8>, bearer: Option<&str>) -> Result<(), ShellError> {
let target = parse_routine_url(url)?;
let client = routine_client()?;
let token = bearer.map(str::trim).filter(|token| !token.is_empty());
let mut builder = client
.post(target)
.header(reqwest::header::CONTENT_TYPE, "application/json");
if let Some(token) = token {
let authorization = bearer_header(token)?;
let mut automation_key =
reqwest::header::HeaderValue::from_str(token).map_err(|_| ShellError::RoutineBearer)?;
automation_key.set_sensitive(true);
builder = builder
.header(reqwest::header::AUTHORIZATION, authorization)
.header("x-automation-key", automation_key);
}
let response = builder.body(bytes).send().map_err(|err| {
ShellError::RoutineTransport(redact_wake(public_reqwest(&err), url, token))
})?;
let status = response.status().as_u16();
if !(200..300).contains(&status) {
return Err(ShellError::RoutineStatus(status));
}
Ok(())
}
fn redact_wake(mut text: String, url: &str, bearer: Option<&str>) -> String {
if !url.is_empty() {
text = text.replace(url, "[redacted]");
}
if let Some(token) = bearer.filter(|token| token.len() >= 4) {
text = text.replace(token, "[redacted]");
}
text
}
fn routine_json(wake: &DecryptedWake<'_>) -> Result<Vec<u8>, ShellError> {
let mut object = serde_json::Map::new();
object.insert(
"body".to_string(),
serde_json::Value::String(wake.body.to_string()),
);
object.insert(
"from".to_string(),
serde_json::Value::String(wake.from.to_string()),
);
object.insert(
"event_id".to_string(),
serde_json::Value::String(wake.event_id.to_string()),
);
if let Some(nick) = wake.nick.map(str::trim).filter(|nick| !nick.is_empty()) {
object.insert(
"nick".to_string(),
serde_json::Value::String(nick.to_string()),
);
}
for (key, value) in [
("room", wake.room),
("from_nick", wake.from_nick),
("to", wake.to),
("reply", wake.reply),
] {
if let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) {
object.insert(
key.to_string(),
serde_json::Value::String(value.to_string()),
);
}
}
serde_json::to_vec(&object)
.map_err(|err| ShellError::RoutineTransport(clip_public(err.to_string())))
}
fn bearer_header(token: &str) -> Result<reqwest::header::HeaderValue, ShellError> {
let token = webhook::bearer_token(token).ok_or(ShellError::RoutineBearer)?;
let mut header = reqwest::header::HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(|_| ShellError::RoutineBearer)?;
header.set_sensitive(true);
Ok(header)
}
fn routine_client() -> Result<reqwest::blocking::Client, ShellError> {
reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(8))
.redirect(reqwest::redirect::Policy::none())
.http1_only()
.use_rustls_tls()
.build()
.map_err(|err| ShellError::RoutineTransport(public_reqwest(&err)))
}
fn public_reqwest(err: &reqwest::Error) -> String {
let mut text = err.to_string();
let mut source = std::error::Error::source(err);
while let Some(inner) = source {
text.push_str(": ");
text.push_str(&inner.to_string());
source = inner.source();
if text.len() > 400 {
break;
}
}
clip_public(text)
}
fn parse_routine_url(url: &str) -> Result<reqwest::Url, ShellError> {
if !webhook::url_ok(url) {
return Err(ShellError::RoutineUrl);
}
let parsed = reqwest::Url::parse(url).map_err(|_| ShellError::RoutineUrl)?;
match parsed.scheme() {
"http" | "https" => {}
_ => return Err(ShellError::RoutineUrl),
}
if parsed.host_str().is_none() || !parsed.username().is_empty() || parsed.password().is_some() {
return Err(ShellError::RoutineUrl);
}
Ok(parsed)
}
pub struct OpenedStore {
dir: PathBuf,
driver: Driver<SealedRecordCodec>,
device_token: Zeroizing<String>,
base_url: String,
session_id: String,
nick: Option<String>,
routine_url: Option<String>,
routine_bearer: Option<Zeroizing<String>>,
leader_sock: Option<PathBuf>,
leader_cwd: Option<String>,
routine_sent: HashSet<String>,
routine_woken_seed_pending: bool,
routine_restart_catchup: bool,
leader_sent: HashSet<String>,
wake_note: Option<String>,
wake_log: Vec<WakeAttempt>,
wake_chain: Option<(provider::ProviderSession, provider::WakeChain)>,
wake_route: Option<String>,
bus: Option<Arc<machine::LocalBus>>,
local_peers: Vec<(String, String)>,
pushed_room_events: Vec<PushedRoomEvent>,
history_pulled: HashSet<String>,
tip_pulled: HashSet<String>,
}
struct ZeroJitter;
impl Jitter for ZeroJitter {
fn next_unit(&mut self) -> f64 {
0.0
}
}
pub struct RoomView {
pub room_id: String,
pub membership: String,
pub encrypted: bool,
}
pub struct TextView {
pub room_id: String,
pub body: String,
pub outcome: String,
pub event_id: Option<String>,
}
impl OpenedStore {
pub fn open(
dir: &Path,
session_id: &str,
device_id: DeviceId,
user_id: &str,
server_name: &str,
backend: Arc<dyn Backend>,
device_token: &str,
) -> Result<Self, ShellError> {
if session_id.is_empty() {
return Err(ShellError::EmptySession);
}
if device_token.is_empty()
|| device_token
.bytes()
.any(|byte| byte.is_ascii_whitespace() || !byte.is_ascii())
{
return Err(ShellError::DeviceToken);
}
let base_url = backend.server_ref().to_string();
fs_create_dir(dir)?;
let user_id = UserId::parse(user_id)?;
let secrets = CoreSecrets {
store_seal_key: Some(store_key::store_key(dir, session_id)?),
backup_key: None,
};
let config = CoreConfig {
user_id,
device_id,
server_name: server_name.to_string(),
};
let records = read_records(dir)?;
let mut core =
match MessengerCore::open_sealed(records, secrets, config, 0, Box::new(ZeroJitter)) {
Ok(core) => core,
Err(MessengerError::Store(err)) => return Err(ShellError::Store(err)),
Err(err) => return Err(ShellError::Messenger(err)),
};
persist(dir, &mut core)?;
let mut opened = Self {
dir: dir.to_path_buf(),
driver: Driver::new(core, backend, {
let dir = dir.to_path_buf();
Box::new(move |core| persist(&dir, core).map_err(|e| AgentError::Store(clip_public(e.to_string()))))
}),
base_url,
device_token: Zeroizing::new(device_token.to_string()),
session_id: session_id.to_string(),
nick: None,
routine_url: None,
routine_bearer: None,
leader_sock: None,
leader_cwd: None,
routine_sent: HashSet::new(),
routine_woken_seed_pending: false,
routine_restart_catchup: false,
leader_sent: HashSet::new(),
wake_chain: None,
wake_route: None,
wake_note: None,
wake_log: Vec::new(),
bus: None,
local_peers: Vec::new(),
pushed_room_events: Vec::new(),
history_pulled: HashSet::new(),
tip_pulled: HashSet::new(),
};
opened.note_already_present();
Ok(opened)
}
pub fn connect(config: &SessionConfig) -> Result<Self, ShellError> {
Self::connect_with_wake(config, SessionWake::default())
}
pub(crate) fn connect_with_wake(config: &SessionConfig, wake: SessionWake) -> Result<Self, ShellError> {
let registered = register_session(config)?;
let server_name = registered
.user_id
.split_once(':')
.map(|(_, server)| server)
.filter(|server| !server.is_empty())
.ok_or(ShellError::Register("user id has no server".to_string()))?;
let mut opened = Self::open(
&config.store_dir(),
&config.session_id,
registered.device_id,
®istered.user_id,
server_name,
registered.backend,
®istered.bearer,
)?;
opened.nick = Some(registered.nick.clone());
opened.set_wake(wake);
opened.drive(1_000, false)?;
Ok(opened)
}
pub fn connect_from_env() -> Result<Self, ShellError> {
Self::connect(&SessionConfig::from_env()?)
}
pub fn connect_node_from_env() -> Result<Self, ShellError> {
let wake = SessionWake::node_cli()?;
Self::connect_with_wake(&SessionConfig::from_env()?, wake)
}
#[cfg(test)]
pub(crate) fn connect_node_from_lookup(
mut get: impl FnMut(&str) -> Option<String>,
toml_text: Option<&str>,
) -> Result<Self, ShellError> {
let wake = SessionWake::node_from_lookup(&mut get)?;
let config = SessionConfig::from_lookup(&mut get, toml_text)?;
Self::connect_with_wake(&config, wake)
}
pub(crate) fn keep_prefix(&self) -> bool {
self.driver.backend().keep_prefix()
}
pub fn device_bearer(&self) -> &str {
self.device_token.as_str()
}
pub fn nick(&self) -> Option<&str> {
self.nick.as_deref()
}
pub fn session_id(&self) -> &str {
&self.session_id
}
pub fn homeserver_url(&self) -> &str {
self.base_url.trim_end_matches('/')
}
pub fn has_leader(&self) -> bool {
self.leader_sock.is_some()
}
pub fn store_dir(&self) -> &Path {
&self.dir
}
pub fn homeserver_hits(&self) -> u64 {
self.bus.as_ref().map(|bus| bus.hits()).unwrap_or(0)
}
pub fn pushed_room_events(&self) -> &[PushedRoomEvent] {
&self.pushed_room_events
}
pub(crate) fn record_push(&mut self, event: PushedRoomEvent) {
self.pushed_room_events.push(event);
}
pub(crate) fn attach_bus(&mut self, bus: Arc<machine::LocalBus>) {
let inner = Arc::clone(self.driver.backend());
self.driver.set_backend(Arc::new(bus_backend::BusBackend { inner, bus: Arc::clone(&bus), user_id: self.driver.core.user_id().as_str().to_string(), force_local: None }));
self.bus = Some(bus);
}
pub(crate) fn set_registered_nick(&mut self, nick: String) {
self.nick = Some(nick);
}
pub(crate) fn set_local_peers(&mut self, peers: Vec<(String, String)>) {
self.local_peers = peers;
}
pub(crate) fn abandon_inflight_sync(&mut self, now_ms: i64) -> Result<(), ShellError> {
if self.driver.sync_inflight() {
self.driver.abandon_sync(now_ms);
self.persist_core()?;
}
Ok(())
}
pub fn user_id(&self) -> &str {
self.driver.core.user_id().as_str()
}
pub fn member_joined(&self, room_id: &str, user_id: &str) -> bool {
let Ok(room_id) = RoomId::parse(room_id) else {
return false;
};
let Ok(user_id) = UserId::parse(user_id) else {
return false;
};
self.driver.core.room_state(&room_id).is_some_and(|state| {
state
.members
.get(&user_id)
.is_some_and(|member| member.membership == Membership::Join)
})
}
pub fn find_nick(
&mut self,
name_or_nick: &str,
now_ms: i64,
) -> Result<FoundSession, ShellError> {
let needle = nick::lookup_nick(name_or_nick)?;
if let Some((nick, user_id)) = self
.local_peers
.iter()
.find(|(nick, _)| nick.eq_ignore_ascii_case(&needle))
{
return Ok(FoundSession {
nick: nick.clone(),
user_id: user_id.clone(),
});
}
self.dispatch(
MessengerCommand::SearchUsers {
term: needle.clone(),
},
now_ms,
)?;
self.drive(now_ms, false)?;
let mut hits: Vec<FoundSession> = self
.driver.core
.user_search_result()
.iter()
.filter_map(|entry| {
let nick = entry.display_name.as_deref()?.trim();
if !nick.eq_ignore_ascii_case(&needle) {
return None;
}
Some(FoundSession {
nick: nick.to_string(),
user_id: entry.user_id.as_str().to_string(),
})
})
.collect();
hits.sort_by(|left, right| left.user_id.cmp(&right.user_id));
hits.dedup_by(|left, right| left.user_id == right.user_id);
match hits.len() {
1 => Ok(hits.remove(0)),
0 => Err(ShellError::UnknownNick),
_ => Err(ShellError::UnknownNick),
}
}
pub fn ensure_dm(&mut self, name_or_nick: &str, mut now_ms: i64) -> Result<String, ShellError> {
let (_peer, room_id) = self.open_dm(name_or_nick, &mut now_ms)?;
Ok(room_id)
}
pub fn accept_direct_invites(&mut self, mut now_ms: i64) -> Result<Vec<String>, ShellError> {
let me = self.driver.core.user_id().clone();
let invited: Vec<RoomId> = self
.driver.core
.room_ids()
.cloned()
.filter(|room_id| {
let Some(state) = self.driver.core.room_state(room_id) else {
return false;
};
state.members.get(&me).is_some_and(|member| {
member.membership == Membership::Invite && member.is_direct
})
})
.collect();
let mut ids = Vec::new();
for room_id in invited {
self.dispatch(
MessengerCommand::JoinRoom {
room_id: room_id.clone(),
},
now_ms,
)?;
ids.push(room_id.as_str().to_string());
}
if ids.is_empty() {
return Ok(ids);
}
for _ in 0..8 {
now_ms += 1_000;
self.drive(now_ms, false)?;
let still_invited = ids.iter().any(|room_id| {
self.rooms()
.iter()
.any(|room| room.room_id == *room_id && room.membership == "invite")
});
if !still_invited {
break;
}
}
Ok(ids)
}
pub fn write_to_nick(
&mut self,
name_or_nick: &str,
text: &str,
mut now_ms: i64,
) -> Result<String, ShellError> {
let (peer, room_id) = self.open_dm(name_or_nick, &mut now_ms)?;
if !self.member_joined(&room_id, peer.as_str()) {
return Err(ShellError::Dm);
}
let encrypted = self
.rooms()
.into_iter()
.any(|room| room.room_id == room_id && room.encrypted);
if !encrypted {
return Err(ShellError::Dm);
}
self.dispatch(
MessengerCommand::SendMessage {
room_id: RoomId::parse(&room_id)?,
message: OutgoingMessage {
kind: MessageKind::Text,
body: text.to_string(),
reply_to: None,
edit_of: None,
},
txn_id: None,
},
now_ms,
)?;
for _ in 0..20 {
now_ms += 1_000;
self.drive(now_ms, false)?;
if let Some(row) = self
.texts()
.into_iter()
.find(|row| row.room_id == room_id && row.body == text)
{
if row.outcome == "sent" {
return Ok(room_id);
}
if row.outcome.starts_with("failed") {
return Err(ShellError::Dm);
}
}
}
Err(ShellError::Dm)
}
fn open_dm(
&mut self,
name_or_nick: &str,
now_ms: &mut i64,
) -> Result<(UserId, String), ShellError> {
let found = self.find_nick(name_or_nick, *now_ms)?;
let peer = UserId::parse(&found.user_id)?;
if &peer == self.driver.core.user_id() {
return Err(ShellError::UnknownNick);
}
let room_id = if let Some(room_id) = self.dm_room(&peer) {
room_id
} else {
self.dispatch(
MessengerCommand::CreateRoom {
kind: CreateRoomKind::Dm { peer: peer.clone() },
},
*now_ms,
)?;
self.wait_for_dm(&peer, now_ms)?
};
Ok((peer, room_id))
}
fn wait_for_dm(&mut self, peer: &UserId, now_ms: &mut i64) -> Result<String, ShellError> {
for attempt in 0..6 {
*now_ms += 1_000;
self.drive(*now_ms, attempt == 5)?;
if let Some(room_id) = self.dm_room(peer) {
return Ok(room_id);
}
}
Err(ShellError::Dm)
}
fn dm_room(&self, peer: &UserId) -> Option<String> {
let me = self.driver.core.user_id().clone();
let ids: Vec<RoomId> = self.driver.core.room_ids().cloned().collect();
for room_id in &ids {
let (joined, has_peer) = {
let Some(state) = self.driver.core.room_state(room_id) else {
continue;
};
let joined = state
.members
.get(&me)
.is_some_and(|member| member.membership == Membership::Join);
let has_peer = state.members.contains_key(peer);
(joined, has_peer)
};
if !joined || !has_peer {
continue;
}
if self.driver.core.room_kind(room_id) == Some(RoomKind::Dm) {
return Some(room_id.as_str().to_string());
}
}
None
}
pub fn set_wake(&mut self, wake: SessionWake) {
self.routine_url = wake.routine_url.filter(|url| !url.is_empty());
self.routine_bearer = wake
.routine_bearer
.filter(|token| !token.is_empty())
.map(Zeroizing::new);
self.leader_sock = wake.leader_sock.filter(|path| !path.as_os_str().is_empty());
self.leader_cwd = wake.leader_cwd.filter(|cwd| !cwd.is_empty());
}
pub fn set_wake_chain(
&mut self,
session: provider::ProviderSession,
chain: provider::WakeChain,
) {
self.wake_chain = Some((session, chain));
}
pub fn wake_chain_ids(&self) -> Vec<&'static str> {
self.wake_chain
.as_ref()
.map(|(_, chain)| chain.ids())
.unwrap_or_default()
}
pub fn wake_route(&self) -> Option<&str> {
self.wake_route.as_deref()
}
pub fn take_wake_route(&mut self) -> Option<String> {
self.wake_route.take()
}
pub fn has_routine(&self) -> bool {
self.routine_url.is_some()
}
pub(crate) fn routine_target(&self) -> Option<(String, Option<String>)> {
let url = self.routine_url.clone()?;
let bearer = self.routine_bearer.as_ref().map(|token| token.as_str().to_string());
Some((url, bearer))
}
pub fn wake_note(&self) -> Option<&str> {
self.wake_note.as_deref()
}
pub fn wake_log(&self) -> &[WakeAttempt] {
&self.wake_log
}
pub fn dispatch(&mut self, command: MessengerCommand, now_ms: i64) -> Result<(), ShellError> {
self.driver.core.dispatch(command, now_ms)?;
self.persist_core()?;
Ok(())
}
pub fn send_room_message(
&mut self,
room_id: &str,
text: &str,
now_ms: i64,
) -> Result<Vec<OutgoingRequest>, ShellError> {
let room_id = RoomId::parse(room_id)?;
self.driver.core.dispatch(
MessengerCommand::SendMessage {
room_id,
message: OutgoingMessage {
kind: MessageKind::Text,
body: text.to_string(),
reply_to: None,
edit_of: None,
},
txn_id: None,
},
now_ms,
)?;
self.persist_core()?;
let mut released = self.driver.core.releasable_requests(now_ms);
self.persist_core()?;
if !released
.iter()
.any(|request| request.kind == OutgoingRequestKind::RoomSend)
{
released.extend(self.driver.core.releasable_requests(now_ms));
self.persist_core()?;
}
Ok(released)
}
pub fn drive(&mut self, now_ms: i64, wait_for_sync: bool) -> Result<(), ShellError> {
self.driver.core.decrypt_loaded_timeline();
let trace_at = self.driver.http_trace.len();
let history = self.request_missing_history(now_ms)?;
self.pull_room_tips()?;
let mut waited_long_poll = false;
if wait_for_sync && self.driver.sync_inflight() {
waited_long_poll = self.driver.harvest_sync(now_ms, true)?;
} else {
self.driver.harvest_sync(now_ms, false)?;
}
self.wake_inbound();
for _ in 0..24 {
if !self.driver.step(now_ms, wait_for_sync, &mut waited_long_poll)? {
break;
}
self.wake_inbound();
}
self.note_history_pages(&history, trace_at);
self.note_room_tips();
self.routine_restart_catchup = false;
Ok(())
}
fn pull_room_tips(&mut self) -> Result<(), ShellError> {
let me = self.driver.core.user_id().clone();
let ids: Vec<RoomId> = self.driver.core.room_ids().cloned().collect();
for room_id in ids {
if self.tip_pulled.contains(room_id.as_str()) {
continue;
}
if self
.driver.core
.timeline(&room_id)
.is_some_and(|timeline| !timeline.items().is_empty())
{
self.tip_pulled.insert(room_id.as_str().to_string());
continue;
}
let joined = self.driver.core.room_state(&room_id).is_some_and(|state| {
state
.members
.get(&me)
.is_some_and(|member| matches!(member.membership, Membership::Join))
});
if !joined {
continue;
}
self.driver.core.pull_latest_page(room_id)?;
}
Ok(())
}
fn note_room_tips(&mut self) {
let ids: Vec<RoomId> = self.driver.core.room_ids().cloned().collect();
for room_id in ids {
if self
.driver.core
.timeline(&room_id)
.is_some_and(|timeline| !timeline.items().is_empty())
{
self.tip_pulled.insert(room_id.as_str().to_string());
}
}
}
fn request_missing_history(&mut self, now_ms: i64) -> Result<Vec<String>, ShellError> {
let me = self.driver.core.user_id().clone();
let ids: Vec<RoomId> = self.driver.core.room_ids().cloned().collect();
let mut requested = Vec::new();
for room_id in ids {
if !self.room_needs_history(&room_id, &me) {
continue;
}
let label = room_id.as_str().to_string();
self.driver.core
.dispatch(MessengerCommand::LoadOlder { room_id }, now_ms)?;
requested.push(label);
}
Ok(requested)
}
fn room_needs_history(&self, room_id: &RoomId, me: &UserId) -> bool {
if self.history_pulled.contains(room_id.as_str()) {
return false;
}
let joined = match self
.driver.core
.room_state(room_id)
.and_then(|state| state.members.get(me))
{
Some(member) => matches!(member.membership, Membership::Join),
None => false,
};
if !joined {
return false;
}
match self.driver.core.timeline(room_id) {
Some(timeline) => timeline.items().is_empty(),
None => true,
}
}
fn note_history_pages(&mut self, requested: &[String], trace_at: usize) {
if requested.is_empty() {
return;
}
let ok = self.driver.http_trace[trace_at..]
.iter()
.filter(|(kind, status)| {
*kind == OutgoingRequestKind::RoomMessages && (200..300).contains(status)
})
.count();
if ok < requested.len() {
return;
}
for room_id in requested {
self.history_pulled.insert(room_id.clone());
}
}
pub fn rooms(&self) -> Vec<RoomView> {
let me = self.driver.core.user_id().clone();
let mut rooms: Vec<RoomView> = self
.driver.core
.room_ids()
.map(|room_id| {
let state = self.driver.core.room_state(room_id);
let membership = state
.and_then(|state| state.members.get(&me))
.map(|member| membership_name(&member.membership).to_string())
.unwrap_or_else(|| "absent".to_string());
let encrypted = state.and_then(|state| state.encryption.as_ref()).is_some();
RoomView {
room_id: room_id.as_str().to_string(),
membership,
encrypted,
}
})
.collect();
rooms.sort_by(|left, right| left.room_id.cmp(&right.room_id));
rooms
}
pub fn texts(&self) -> Vec<TextView> {
let mut out = Vec::new();
for room_id in self.driver.core.room_ids() {
let Some(timeline) = self.driver.core.timeline(room_id) else {
continue;
};
for item in timeline.items() {
match &item.content {
ItemContent::Text(text)
| ItemContent::Notice(text)
| ItemContent::Emote(text) => {
out.push(TextView {
room_id: room_id.as_str().to_string(),
body: text.body.clone(),
outcome: outcome_name(&item.send_state),
event_id: item.event_id.as_ref().map(|id| id.as_str().to_string()),
});
}
ItemContent::Undecryptable { reason } => out.push(TextView {
room_id: room_id.as_str().to_string(),
body: String::new(),
outcome: format!("undecryptable:{reason}"),
event_id: item.event_id.as_ref().map(|id| id.as_str().to_string()),
}),
_ => {}
}
}
}
out
}
pub fn sync_inflight(&self) -> bool {
self.driver.sync_inflight()
}
pub fn http_trace(&self) -> Vec<String> {
self.driver.http_trace
.iter()
.map(|(kind, status)| format!("{kind:?} {status}"))
.collect()
}
pub fn take_ingest_error(&mut self) -> Option<String> {
self.driver.core.take_ingest_error().map(clip_public)
}
fn persist_core(&mut self) -> Result<(), ShellError> {
self.driver.persist().map_err(ShellError::from)
}
pub fn take_security_alerts(&mut self) -> Vec<String> {
std::mem::take(&mut self.driver.security_alerts)
}
fn inbound_plaintexts(&self) -> Vec<InboundPlaintext> {
let me = self.driver.core.user_id();
let mut out = Vec::new();
for room_id in self.driver.core.room_ids() {
let Some(timeline) = self.driver.core.timeline(room_id) else {
continue;
};
for item in timeline.items() {
if item.redacted || &item.sender == me || item.send_state != SendState::Sent {
continue;
}
let body = match &item.content {
ItemContent::Text(text)
| ItemContent::Notice(text)
| ItemContent::Emote(text) => text.body.clone(),
_ => continue,
};
let Some(event_id) = item.event_id.as_ref() else {
continue;
};
let key = format!("{}\n{}", room_id.as_str(), event_id.as_str());
out.push(InboundPlaintext {
room_id: room_id.as_str().to_string(),
key,
body,
from: item.sender.as_str().to_string(),
nick: self.sender_nick(room_id, &item.sender),
event_id: event_id.as_str().to_string(),
});
}
}
out
}
fn sender_nick(&self, room_id: &RoomId, sender: &UserId) -> Option<String> {
let name = self
.driver.core
.room_state(room_id)?
.members
.get(sender)?
.displayname
.as_deref()?
.trim();
if name.is_empty() {
None
} else {
Some(name.to_string())
}
}
fn note_already_present(&mut self) {
self.load_routine_woken();
self.load_leader_prompted();
}
fn load_routine_woken(&mut self) {
let path = self.dir.join(ROUTINE_WOKEN_FILE);
match std::fs::read_to_string(&path) {
Ok(text) if !text.trim().is_empty() => {
for line in text.lines() {
let Some((room, event)) = line.split_once('\t') else {
continue;
};
if room.is_empty() || event.is_empty() {
continue;
}
self.routine_sent.insert(format!("{room}\n{event}"));
}
self.routine_woken_seed_pending = false;
self.routine_restart_catchup = true;
}
Ok(_) | Err(_) => {
self.routine_woken_seed_pending = true;
self.routine_restart_catchup = false;
}
}
}
fn seed_routine_woken_from_timeline(&mut self) {
let path = self.dir.join(ROUTINE_WOKEN_FILE);
let items = self.inbound_plaintexts();
let mut newest: std::collections::HashMap<String, String> =
std::collections::HashMap::new();
for item in &items {
newest.insert(item.room_id.clone(), item.key.clone());
}
let mut lines = String::new();
for item in &items {
if newest.get(&item.room_id) == Some(&item.key) {
continue;
}
self.routine_sent.insert(item.key.clone());
if let Some((room, event)) = item.key.split_once('\n') {
lines.push_str(room);
lines.push('\t');
lines.push_str(event);
lines.push('\n');
}
}
let _ = std::fs::write(path, lines);
self.routine_woken_seed_pending = false;
}
fn remember_routine_wake(&mut self, key: &str) {
self.routine_sent.insert(key.to_string());
let Some((room, event)) = key.split_once('\n') else {
return;
};
if room.is_empty() || event.is_empty() {
return;
}
let path = self.dir.join(ROUTINE_WOKEN_FILE);
let mut file = match std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&path)
{
Ok(file) => file,
Err(_) => return,
};
use std::io::Write;
let _ = write!(file, "{room}\t{event}\n");
}
fn load_leader_prompted(&mut self) {
let path = self.dir.join(LEADER_PROMPTED_FILE);
match std::fs::read_to_string(&path) {
Ok(text) => {
for line in text.lines() {
let Some((room, event)) = line.split_once('\t') else {
continue;
};
if room.is_empty() || event.is_empty() {
continue;
}
self.leader_sent.insert(format!("{room}\n{event}"));
}
}
Err(_) => self.seed_leader_prompted_except_newest(&path),
}
}
fn seed_leader_prompted_except_newest(&mut self, path: &Path) {
let items = self.inbound_plaintexts();
let mut newest: std::collections::HashMap<String, String> =
std::collections::HashMap::new();
for item in &items {
newest.insert(item.room_id.clone(), item.key.clone());
}
let mut lines = String::new();
for item in &items {
if newest.get(&item.room_id) == Some(&item.key) {
continue;
}
self.leader_sent.insert(item.key.clone());
if let Some((room, event)) = item.key.split_once('\n') {
lines.push_str(room);
lines.push('\t');
lines.push_str(event);
lines.push('\n');
}
}
let _ = std::fs::write(path, lines);
}
fn remember_leader_prompt(&mut self, key: &str) {
self.leader_sent.insert(key.to_string());
let Some((room, event)) = key.split_once('\n') else {
return;
};
if room.is_empty() || event.is_empty() {
return;
}
let path = self.dir.join(LEADER_PROMPTED_FILE);
let mut file = match std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&path)
{
Ok(file) => file,
Err(_) => return,
};
use std::io::Write;
let _ = write!(file, "{room}\t{event}\n");
}
fn arm_unprompted_newest(&mut self) {
if (self.leader_sock.is_none() && self.wake_chain.is_none()) || !self.leader_sent.is_empty() {
return;
}
if self.inbound_plaintexts().is_empty() {
return;
}
let path = self.dir.join(LEADER_PROMPTED_FILE);
self.seed_leader_prompted_except_newest(&path);
}
#[cfg_attr(not(feature = "wake-grok"), allow(unused_variables))]
fn wake_inbound(&mut self) {
if self.routine_url.is_none() && self.leader_sock.is_none() && self.wake_chain.is_none() {
return;
}
if self.routine_restart_catchup && self.routine_url.is_some() {
for item in self.inbound_plaintexts() {
if !self.routine_sent.contains(&item.key) {
self.remember_routine_wake(&item.key);
}
}
} else if self.routine_woken_seed_pending && self.routine_url.is_some() {
if !self.inbound_plaintexts().is_empty() {
self.seed_routine_woken_from_timeline();
}
}
self.arm_unprompted_newest();
let url = self.routine_url.clone();
let bearer = self.routine_bearer.clone();
let sock = self.leader_sock.clone();
let cwd = self.leader_cwd.clone();
let session_id = self.session_id.clone();
let items = self.inbound_plaintexts();
for item in items {
if let Some(url) = url.as_deref() {
if !self.routine_restart_catchup && !self.routine_sent.contains(&item.key) {
let from_nick = mxid_localpart(&item.from).to_string();
let to = self.nick.clone();
let reply = to.as_deref().map(|to| reply_hint(to, &from_nick));
let wake = DecryptedWake {
body: &item.body,
from: &item.from,
nick: item.nick.as_deref(),
event_id: &item.event_id,
room: Some(&item.room_id),
from_nick: Some(&from_nick),
to: to.as_deref(),
reply: reply.as_deref(),
};
match post_decrypted_with_bearer(
url,
&wake,
bearer.as_ref().map(|token| token.as_str()),
) {
Ok(()) => {
self.remember_routine_wake(&item.key);
self.wake_log.push(WakeAttempt {
event_id: item.event_id.clone(),
status: Some(200),
});
}
Err(err) => {
let status = match &err {
ShellError::RoutineStatus(code) => Some(*code),
_ => None,
};
self.wake_log.push(WakeAttempt {
event_id: item.event_id.clone(),
status,
});
self.wake_note = Some(clip_public(err.to_string()));
}
}
}
}
if let Some(sock) = sock.as_deref() {
if !self.leader_sent.contains(&item.key) {
let cwd = cwd.clone().or_else(|| {
std::env::current_dir()
.ok()
.map(|path| path.display().to_string())
});
let Some(cwd) = cwd.filter(|cwd| !cwd.is_empty()) else {
self.wake_note = Some("leader cwd is empty".to_string());
continue;
};
#[cfg(feature = "wake-grok")]
{
match mail4agent_grok::wake_decrypted_room_blocking(
sock,
&session_id,
&cwd,
&item.body,
) {
Ok(()) => {
self.remember_leader_prompt(&item.key);
}
Err(err) => {
self.wake_note = Some(clip_public(err.to_string()));
}
}
#[cfg(not(feature = "wake-grok"))]
{
let _ = (sock, &session_id, &cwd);
self.wake_note = Some("this build was made without feature wake-grok".to_string());
}
}
}
}
if sock.is_none() && !self.leader_sent.contains(&item.key) {
let Some((session, mut chain)) = self.wake_chain.take() else {
continue;
};
let from_nick = mxid_localpart(&item.from).to_string();
let letter = provider::WakeLetter {
body: &item.body,
from_nick: &from_nick,
event_id: &item.event_id,
room: Some(&item.room_id),
};
match chain.wake(&session, &letter) {
Ok((id, outcome)) => {
self.remember_leader_prompt(&item.key);
let how = match outcome {
provider::WakeOutcome::Delivered => "delivered",
provider::WakeOutcome::Queued(_) => "queued",
};
self.wake_route = Some(format!("{id} {how}"));
}
Err(err) => {
self.wake_note = Some(clip_public(err.to_string()));
}
}
self.wake_chain = Some((session, chain));
}
}
}
}
struct InboundPlaintext {
room_id: String,
key: String,
body: String,
from: String,
nick: Option<String>,
event_id: String,
}
fn membership_name(membership: &Membership) -> &'static str {
match membership {
Membership::Invite => "invite",
Membership::Join => "join",
Membership::Knock => "knock",
Membership::Leave => "leave",
Membership::Ban => "ban",
Membership::Unknown => "unknown",
}
}
fn outcome_name(state: &SendState) -> String {
match state {
SendState::Sent => "sent".to_string(),
SendState::Sending | SendState::LocalEcho => "sending".to_string(),
SendState::Failed { reason } => format!("failed:{reason}"),
}
}
fn parse_base_url(raw: &str) -> Result<reqwest::Url, ShellError> {
let url = reqwest::Url::parse(raw).map_err(|_| ShellError::BaseUrl)?;
match url.scheme() {
"http" | "https" => {}
_ => return Err(ShellError::BaseUrl),
}
if url.host_str().is_none()
|| url.query().is_some()
|| url.fragment().is_some()
|| !url.username().is_empty()
|| url.password().is_some()
{
return Err(ShellError::BaseUrl);
}
Ok(url)
}
struct RegisteredSession {
user_id: String,
device_id: DeviceId,
bearer: Zeroizing<String>,
nick: String,
backend: Arc<dyn Backend>,
}
#[cfg(any(feature = "tier-server", feature = "tier-matrix"))]
fn identity_session(config: &SessionConfig, auth: &IdentityAuth) -> Result<RegisteredSession, ShellError> {
let fail = |e: AgentError| ShellError::Register(clip_public(e.to_string()));
parse_base_url(&config.homeserver_url)?;
let ids = m4a_agent::IdentityStore::new(store_key::vault(&config.store_root)?);
let backend: Arc<dyn Backend> = match auth.tier {
#[cfg(feature = "tier-server")]
m4a_agent::BackendKind::Server => Arc::new(m4a_agent::backend::server::ServerBackend::new(&config.homeserver_url).map_err(fail)?),
#[cfg(feature = "tier-matrix")]
m4a_agent::BackendKind::Matrix => Arc::new(m4a_agent::backend::matrix::MatrixBackend::discover(&config.homeserver_url).map_err(fail)?),
#[allow(unreachable_patterns)]
_ => return Err(ShellError::Register("this build does not include that tier".into())),
};
let mut id = ids.resolve(&config.session_id, auth.tier, backend.server_ref()).map_err(fail)?;
let session = backend.ensure_session(&ids, &mut id, auth.invite.as_ref().map(|i| i.as_str())).map_err(fail)?;
let _ = std::fs::remove_file(invite_path(&config.store_root, &config.session_id));
ids.adopt_nick(&mut id, &session.nick).map_err(fail)?;
let (user_id, device_raw) = match (&session.user_id, &session.device_id) {
(Some(u), Some(d)) => (u.clone(), d.clone()),
_ => {
backend.whoami().map_err(fail)?
}
};
if user_id.is_empty() || device_raw.is_empty() {
return Err(ShellError::Register("whoami missed an id".into()));
}
Ok(RegisteredSession { user_id, device_id: DeviceId::parse(&device_raw)?, bearer: session.token.clone(), nick: session.nick, backend })
}
#[cfg(not(any(feature = "tier-server", feature = "tier-matrix")))]
fn identity_session(_: &SessionConfig, _: &IdentityAuth) -> Result<RegisteredSession, ShellError> {
Err(ShellError::Register("this build has no tier for identity login".into()))
}
fn register_session(config: &SessionConfig) -> Result<RegisteredSession, ShellError> {
identity_session(config, &config.identity)
}
fn clip_public(text: String) -> String {
let mut out = String::new();
for token in text.split_whitespace() {
if token.len() > 80 {
out.push_str("[omitted]");
} else {
out.push_str(token);
}
out.push(' ');
if out.len() > 240 {
break;
}
}
out
}
fn fs_create_dir(dir: &Path) -> Result<(), ShellError> {
std::fs::create_dir_all(dir)?;
Ok(())
}
#[derive(Debug, thiserror::Error)]
pub enum ShellError {
#[error("session id is empty")]
EmptySession,
#[error("store root is unset")]
StoreRoot,
#[error("homeserver url is unset")]
HomeserverUrl,
#[error("bot display name is unset")]
BotName,
#[error("nick could not be derived from the bot display name")]
Nick,
#[error("session list: {0}")]
SessionList(String),
#[error("homeserver register failed: {0}")]
Register(String),
#[error("no session with that nick")]
UnknownNick,
#[error("direct room was not ready")]
Dm,
#[error("config: {0}")]
Config(String),
#[error("device token is empty or not a single header value")]
DeviceToken,
#[error("base url must be http or https")]
BaseUrl,
#[error("record key escapes the store directory")]
BadRecordKey,
#[error("routine url is not an http or https url")]
RoutineUrl,
#[error("node cli does not take a routine url")]
NodeRoutine,
#[error("gateway: {0}")]
Gateway(String),
#[error("routine bearer is empty or not a single header value")]
RoutineBearer,
#[error("routine status {0}")]
RoutineStatus(u16),
#[error("routine post failed: {0}")]
RoutineTransport(String),
#[error("homeserver http failed: {0}")]
Http(String),
#[error("homeserver response was not ingested: {0}")]
Ingest(String),
#[error("io: {0}")]
Io(#[from] std::io::Error),
#[error(transparent)]
Store(#[from] StoreError),
#[error(transparent)]
Messenger(#[from] MessengerError),
}
impl From<AgentError> for ShellError {
fn from(err: AgentError) -> Self {
match err {
AgentError::Transport(text) => ShellError::Http(text),
AgentError::Ingest(text) => ShellError::Ingest(text),
other => ShellError::Register(clip_public(other.to_string())),
}
}
}
fn persist(dir: &Path, core: &mut MessengerCore<SealedRecordCodec>) -> Result<(), ShellError> {
while let Some(batch) = core.take_flush_batch() {
for record in &batch.records {
let path = record_path(dir, &record.key)?;
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(&path, &record.bytes)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600))?;
}
}
for key in &batch.deletes {
let path = record_path(dir, key)?;
if path.exists() {
std::fs::remove_file(path)?;
}
}
core.ack_flush(batch.id);
}
Ok(())
}
const NOT_RECORDS: &[&str] = &[
machine::WAKE_KEYCHAIN_FILE,
machine::STORE_LOCK_FILE,
LEADER_PROMPTED_FILE,
ROUTINE_WOKEN_FILE,
];
const LEADER_PROMPTED_FILE: &str = "leader-prompted";
const ROUTINE_WOKEN_FILE: &str = "routine-woken";
fn read_records(dir: &Path) -> Result<Vec<SealedRecord>, ShellError> {
let mut records = Vec::new();
if dir.exists() {
walk(dir, dir, &mut records)?;
}
Ok(records)
}
fn walk(dir: &Path, root: &Path, records: &mut Vec<SealedRecord>) -> Result<(), ShellError> {
for entry in std::fs::read_dir(dir)? {
let entry = entry?;
let path = entry.path();
if path.is_dir() {
walk(&path, root, records)?;
continue;
}
if dir == root
&& path
.file_name()
.and_then(|name| name.to_str())
.is_some_and(|name| NOT_RECORDS.contains(&name))
{
continue;
}
let rel = path
.strip_prefix(root)
.map_err(|_| ShellError::BadRecordKey)?;
let key = rel
.components()
.map(|component| component.as_os_str().to_string_lossy())
.collect::<Vec<_>>()
.join("/");
let bytes = std::fs::read(&path)?;
records.push(SealedRecord {
key: RecordKey::new(key),
bytes,
});
}
Ok(())
}
fn record_path(dir: &Path, key: &RecordKey) -> Result<PathBuf, ShellError> {
let rel = Path::new(key.as_str());
if rel.is_absolute()
|| rel.components().any(|component| {
matches!(
component,
Component::ParentDir | Component::RootDir | Component::Prefix(_)
)
})
{
return Err(ShellError::BadRecordKey);
}
Ok(dir.join(rel))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn session_config_is_identity_only_and_reads_no_secret() {
let get = |extra: &'static [(&'static str, &'static str)]| {
move |key: &str| {
extra.iter().find(|(k, _)| *k == key).map(|(_, v)| v.to_string()).or_else(|| match key {
SESSION_ID_ENV => Some("web-session-1".to_string()),
STORE_ROOT_ENV => Some("/tmp/m4a-root".to_string()),
"M4A_PRODUCT_PASSWORD" | "M4A_PRODUCT_TOKEN" | "M4A_DEVICE_TOKEN" => Some("must-not-be-read".to_string()),
_ => None,
})
}
};
let config = SessionConfig::from_lookup(get(&[(HOMESERVER_URL_ENV, "http://127.0.0.1:9")]), None).expect("config");
assert_eq!(config.store_dir(), session_store_dir(Path::new("/tmp/m4a-root"), "web-session-1"));
assert!(!format!("{config:?}").contains("must-not-be-read"));
assert_eq!(config.identity.tier, m4a_agent::BackendKind::Server);
let matrix = SessionConfig::from_lookup(get(&[(PRODUCT_URL_ENV, "http://127.0.0.1:8"), (TIER_ENV, "matrix")]), None).expect("matrix");
assert_eq!(matrix.identity.tier, m4a_agent::BackendKind::Matrix);
assert!(SessionConfig::from_lookup(get(&[(TIER_ENV, "other"), (PRODUCT_URL_ENV, "http://127.0.0.1:8")]), None).is_err());
let from_toml = SessionConfig::from_lookup(get(&[]), Some("homeserver_url = \"http://127.0.0.1:9\"\nother = \"ignored\"\n")).expect("toml url");
assert_eq!(from_toml.homeserver_url, "http://127.0.0.1:9");
assert!(SessionConfig::from_lookup(get(&[]), None).is_err(), "no server given");
}
use std::io::{Read, Write};
use std::net::TcpListener;
use std::process::{Command, Stdio};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::Instant;
fn plain(base: &str, token: &str) -> Arc<dyn Backend> {
Arc::new(m4a_agent::backend::attached::AttachedBackend::with_prefix(m4a_agent::BackendKind::Server, base, token, false).expect("backend"))
}
struct TempDir(PathBuf);
impl Drop for TempDir {
fn drop(&mut self) {
let _ = std::fs::remove_dir_all(&self.0);
}
}
fn temp_dir(name: &str) -> TempDir {
let dir = TempDir(std::env::temp_dir().join(format!(
"mail4agent-messenger-shell-{}-{name}",
std::process::id()
)));
let _ = std::fs::remove_dir_all(&dir.0);
dir
}
fn open_alice(dir: &Path, session: &str) -> OpenedStore {
let device = DeviceId::parse("DEVICE1").expect("device id");
OpenedStore::open(
dir,
session,
device,
"@alice:localhost",
"localhost",
plain("http://127.0.0.1:9", "fake-token"),
"fake-token",
)
.expect("open")
}
#[test]
fn same_session_opens_and_a_different_session_fails_the_seal() {
let dir = temp_dir("seal");
open_alice(&dir.0, "session-a");
open_alice(&dir.0, "session-a");
let device = DeviceId::parse("DEVICE1").expect("device id");
match OpenedStore::open(
&dir.0,
"session-b",
device,
"@alice:localhost",
"localhost",
plain("http://127.0.0.1:9", "fake-token"),
"fake-token",
) {
Err(ShellError::Store(StoreError::CodecOpen { .. })) => {}
Ok(_) => panic!("different session opened the sealed store"),
Err(err) => panic!("expected seal auth failure, got {err}"),
}
}
#[test]
fn two_sessions_under_one_root_seal_and_reopen_only_with_their_own_id() {
let root = temp_dir("isolate");
let dir_a = session_store_dir(&root.0, "session-a");
let dir_b = session_store_dir(&root.0, "session-b");
assert_ne!(dir_a, dir_b);
assert_eq!(dir_a.parent(), Some(root.0.as_path()));
assert_eq!(dir_b.parent(), Some(root.0.as_path()));
assert_eq!(session_store_dir(&root.0, "session-a"), dir_a);
let slipped = session_store_dir(&root.0, "../session-a");
assert_eq!(slipped.parent(), Some(root.0.as_path()));
let name = slipped.file_name().expect("name").to_string_lossy();
assert_eq!(name.len(), 64);
assert!(name.chars().all(|ch| ch.is_ascii_hexdigit()));
let bearer = "isolation-bearer-7c2e";
let device = DeviceId::parse("DEVICE1").expect("device id");
let open_with = |dir: &Path, session: &str| {
OpenedStore::open(
dir,
session,
device.clone(),
"@alice:localhost",
"localhost",
plain("http://127.0.0.1:9", bearer),
bearer,
)
};
open_with(&dir_a, "session-a").expect("seal a");
open_with(&dir_b, "session-b").expect("seal b");
open_with(&dir_a, "session-a").expect("reopen a");
open_with(&dir_b, "session-b").expect("reopen b");
for (dir, session) in [(&dir_a, "session-b"), (&dir_b, "session-a")] {
match open_with(dir, session) {
Err(ShellError::Store(StoreError::CodecOpen { .. })) => {}
Ok(_) => panic!("{session} opened the other sealed store"),
Err(err) => panic!("expected seal auth failure, got {err}"),
}
}
fn contains_bearer(dir: &Path, needle: &str) -> bool {
if !dir.exists() {
return false;
}
for entry in std::fs::read_dir(dir).expect("read") {
let entry = entry.expect("entry");
let path = entry.path();
if path.to_string_lossy().contains(needle) {
return true;
}
if path.is_dir() {
if contains_bearer(&path, needle) {
return true;
}
} else if std::fs::read(&path)
.expect("file")
.windows(needle.len())
.any(|window| window == needle.as_bytes())
{
return true;
}
}
false
}
assert!(
!contains_bearer(&root.0, bearer),
"raw bearer was written under the store root"
);
}
#[test]
fn post_decrypted_reaches_the_loopback_routine() {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let server = thread::spawn(move || {
let (mut sock, _) = listener.accept().expect("accept");
sock.set_read_timeout(Some(Duration::from_secs(2)))
.expect("timeout");
let mut buf = Vec::new();
let mut tmp = [0u8; 1024];
loop {
let n = sock.read(&mut tmp).unwrap_or(0);
if n == 0 {
break;
}
buf.extend_from_slice(&tmp[..n]);
if let Some(header_end) = buf.windows(4).position(|window| window == b"\r\n\r\n") {
let headers = String::from_utf8_lossy(&buf[..header_end]).to_string();
let length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
if name.eq_ignore_ascii_case("content-length") {
value.trim().parse::<usize>().ok()
} else {
None
}
})
.unwrap_or(0);
if buf.len() >= header_end + 4 + length {
let body = buf[header_end + 4..header_end + 4 + length].to_vec();
let _ = sock.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
);
return (headers, body);
}
}
}
(String::new(), buf)
});
let wake = RoutineWake {
url: format!("http://{addr}/routine"),
};
post_decrypted(
&wake.url,
&DecryptedWake {
body: "hello-from-room",
from: "@bob:localhost",
nick: None,
event_id: "$m1:localhost",
..Default::default()
},
)
.expect("post");
let (headers, body) = server.join().expect("listener stopped");
assert!(headers.starts_with("POST /routine HTTP/1.1"), "{headers}");
assert!(
headers
.to_ascii_lowercase()
.contains("content-type: application/json"),
"{headers}"
);
assert!(
!headers.to_ascii_lowercase().contains("authorization"),
"routine post must not add a bearer"
);
assert!(
!headers.to_ascii_lowercase().contains("x-automation-key"),
"routine post must not add a key without a bearer"
);
assert!(!headers.contains("/mail/send"), "{headers}");
assert_wake_json(
&body,
"hello-from-room",
"@bob:localhost",
"$m1:localhost",
None,
);
}
#[test]
fn post_decrypted_attempts_https_against_a_local_self_signed_listener() {
let dir = temp_dir("https");
std::fs::create_dir_all(&dir.0).expect("dir");
let cert = dir.0.join("cert.pem");
let key = dir.0.join("key.pem");
let generated = Command::new("openssl")
.args([
"req",
"-x509",
"-newkey",
"rsa:2048",
"-keyout",
key.to_str().expect("utf-8"),
"-out",
cert.to_str().expect("utf-8"),
"-days",
"1",
"-nodes",
"-subj",
"/CN=127.0.0.1",
])
.stdout(Stdio::null())
.stderr(Stdio::null())
.status()
.expect("openssl req");
assert!(generated.success(), "openssl did not write a local cert");
let probe = TcpListener::bind("127.0.0.1:0").expect("bind");
let port = probe.local_addr().expect("addr").port();
drop(probe);
let mut server = Command::new("openssl")
.args([
"s_server",
"-accept",
&format!("127.0.0.1:{port}"),
"-cert",
cert.to_str().expect("utf-8"),
"-key",
key.to_str().expect("utf-8"),
"-www",
])
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.expect("openssl s_server");
let started = std::time::Instant::now();
while started.elapsed() < Duration::from_secs(3) {
if std::net::TcpStream::connect(("127.0.0.1", port)).is_ok() {
break;
}
thread::sleep(Duration::from_millis(20));
}
let err = post_decrypted(
&format!("https://127.0.0.1:{port}/hook"),
&DecryptedWake {
body: "hello-https",
from: "@bob:localhost",
nick: None,
event_id: "$m1:localhost",
..Default::default()
},
)
.expect_err("a self-signed cert must not verify");
let _ = server.kill();
let _ = server.wait();
assert!(
!matches!(err, ShellError::RoutineUrl),
"https was refused before the client tried it: {err}"
);
let text = err.to_string().to_ascii_lowercase();
assert!(
text.contains("cert")
|| text.contains("tls")
|| text.contains("handshake")
|| text.contains("ssl")
|| text.contains("unknownissuer")
|| text.contains("invalidpeer"),
"https did not reach a tls failure: {err}"
);
}
#[test]
fn a_base_url_carries_no_fragment_credentials_or_query() {
assert!(parse_base_url("http://127.0.0.1:9").is_ok());
assert!(parse_base_url("http://127.0.0.1:9/#other").is_err());
assert!(parse_base_url("http://u:p@127.0.0.1:9").is_err());
assert!(parse_base_url("ftp://127.0.0.1:9").is_err());
}
#[test]
fn send_message_releases_a_matrix_room_send() {
let dir = temp_dir("send");
let mut store = open_alice(&dir.0, "session-a");
let released = store
.send_room_message("!room:localhost", "hello room", 0)
.expect("release");
let send = released
.iter()
.find(|request| request.kind == OutgoingRequestKind::RoomSend)
.unwrap_or_else(|| {
panic!(
"engine did not release a RoomSend: {:?}",
released
.iter()
.map(|request| request.kind)
.collect::<Vec<_>>()
)
});
assert!(
send.path.starts_with("/_matrix/client/v3/rooms/")
&& send.path.contains("/send/m.room.message/"),
"not a matrix room send: {}",
send.path
);
assert!(!send.path.contains("/mail/send"), "{}", send.path);
assert!(!send.path.contains("/admin/listener"), "{}", send.path);
let body = send.body.as_ref().expect("room send body");
assert_eq!(body["msgtype"], "m.text");
assert_eq!(body["body"], "hello room");
}
#[test]
fn drive_performs_the_released_request_with_bearer_and_no_matrix_prefix() {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let seen = Arc::new(Mutex::new(Vec::<String>::new()));
let seen_worker = Arc::clone(&seen);
listener.set_nonblocking(true).expect("nonblocking");
let server = thread::spawn(move || {
let deadline = std::time::Instant::now() + Duration::from_secs(3);
loop {
if std::time::Instant::now() > deadline {
break;
}
let sock = match listener.accept() {
Ok((sock, _)) => sock,
Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(20));
continue;
}
Err(_) => break,
};
let mut sock = sock;
let _ = sock.set_nonblocking(false);
let _ = sock.set_read_timeout(Some(Duration::from_secs(2)));
let mut buf = Vec::new();
let mut tmp = [0u8; 2048];
loop {
let n = sock.read(&mut tmp).unwrap_or(0);
if n == 0 {
break;
}
buf.extend_from_slice(&tmp[..n]);
if let Some(header_end) =
buf.windows(4).position(|window| window == b"\r\n\r\n")
{
let headers = String::from_utf8_lossy(&buf[..header_end]).to_string();
let length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
if name.eq_ignore_ascii_case("content-length") {
value.trim().parse::<usize>().ok()
} else {
None
}
})
.unwrap_or(0);
if buf.len() >= header_end + 4 + length {
let body = buf[header_end + 4..header_end + 4 + length].to_vec();
let first = headers.lines().next().unwrap_or("").to_string();
let bearer_ok = headers.lines().any(|line| {
let (name, value) = line.split_once(':').unwrap_or(("", ""));
name.eq_ignore_ascii_case("authorization")
&& value.trim() == "Bearer fake-token"
});
let note = format!(
"{first} bearer_ok={bearer_ok} body={}",
String::from_utf8_lossy(&body)
);
seen_worker.lock().expect("seen").push(note);
let request = first;
let response_body = if request.contains("/send/") {
br#"{"event_id":"$e:localhost"}"#.to_vec()
} else if request.contains("/keys/") {
br#"{"one_time_key_counts":{"signed_curve25519":50}}"#.to_vec()
} else {
br#"{"next_batch":"s1","device_one_time_keys_count":{"signed_curve25519":50},"device_unused_fallback_key_types":["signed_curve25519"]}"#.to_vec()
};
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
response_body.len()
);
let _ = sock.write_all(head.as_bytes());
let _ = sock.write_all(&response_body);
break;
}
}
}
}
});
let dir = temp_dir("drive");
let device = DeviceId::parse("DEVICE1").expect("device id");
let mut store = OpenedStore::open(
&dir.0,
"session-a",
device,
"@alice:localhost",
"localhost",
plain(&format!("http://{addr}"), "fake-token"),
"fake-token",
)
.expect("open");
store
.dispatch(
MessengerCommand::SendMessage {
room_id: RoomId::parse("!room:localhost").expect("room"),
message: OutgoingMessage {
kind: MessageKind::Text,
body: "hello room".to_string(),
reply_to: None,
edit_of: None,
},
txn_id: None,
},
0,
)
.expect("dispatch");
store.drive(0, false).expect("drive");
thread::sleep(Duration::from_millis(200));
drop(store);
let _ = server.join();
let seen = seen.lock().expect("seen");
assert!(
seen.iter().any(|line| {
line.contains("PUT /client/v3/rooms/")
&& line.contains("/send/m.room.message/")
&& line.contains("bearer_ok=true")
&& line.contains("hello room")
&& !line.contains("/_matrix")
}),
"released room send was not performed: {seen:?}"
);
assert!(
seen.iter().all(|line| !line.contains("/_matrix")),
"prefix was not stripped: {seen:?}"
);
}
struct Hit {
headers: String,
body: Vec<u8>,
}
fn read_http(sock: &mut std::net::TcpStream) -> Option<(String, Vec<u8>)> {
let _ = sock.set_read_timeout(Some(Duration::from_secs(2)));
let mut buf = Vec::new();
let mut tmp = [0u8; 2048];
loop {
let n = sock.read(&mut tmp).unwrap_or(0);
if n == 0 {
break;
}
buf.extend_from_slice(&tmp[..n]);
if let Some(header_end) = buf.windows(4).position(|window| window == b"\r\n\r\n") {
let headers = String::from_utf8_lossy(&buf[..header_end]).to_string();
let length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
if name.eq_ignore_ascii_case("content-length") {
value.trim().parse::<usize>().ok()
} else {
None
}
})
.unwrap_or(0);
if buf.len() >= header_end + 4 + length {
let body = buf[header_end + 4..header_end + 4 + length].to_vec();
return Some((headers, body));
}
}
}
None
}
fn write_http(sock: &mut std::net::TcpStream, status: &str, body: &[u8]) {
let head = format!(
"HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
let _ = sock.write_all(head.as_bytes());
let _ = sock.write_all(body);
}
fn assert_wake_json(bytes: &[u8], body: &str, from: &str, event_id: &str, nick: Option<&str>) {
let parsed: serde_json::Value = serde_json::from_slice(bytes).expect("json wake");
assert_eq!(parsed["body"], body);
assert_eq!(parsed["from"], from);
assert_eq!(parsed["event_id"], event_id);
match nick {
Some(nick) => assert_eq!(parsed["nick"], nick),
None => assert!(
parsed.get("nick").is_none() || parsed["nick"].is_null(),
"unknown nick must be omitted or null, not invented: {parsed}"
),
}
}
fn sync_with_text_nick(body: &str, nick: Option<&str>) -> Vec<u8> {
let mut state = Vec::new();
if let Some(nick) = nick {
state.push(serde_json::json!({
"event_id": "$mem:localhost",
"type": "m.room.member",
"state_key": "@bob:localhost",
"sender": "@bob:localhost",
"origin_server_ts": 1,
"content": { "membership": "join", "displayname": nick }
}));
}
serde_json::json!({
"next_batch": "s1",
"rooms": {
"join": {
"!r:localhost": {
"state": { "events": state },
"timeline": {
"events": [{
"event_id": "$m1:localhost",
"type": "m.room.message",
"sender": "@bob:localhost",
"origin_server_ts": 10,
"content": { "msgtype": "m.text", "body": body }
}]
}
}
}
},
"device_one_time_keys_count": { "signed_curve25519": 50 },
"device_unused_fallback_key_types": ["signed_curve25519"]
})
.to_string()
.into_bytes()
}
fn sync_empty() -> Vec<u8> {
serde_json::json!({
"next_batch": "s2",
"device_one_time_keys_count": { "signed_curve25519": 50 },
"device_unused_fallback_key_types": ["signed_curve25519"]
})
.to_string()
.into_bytes()
}
fn spawn_homeserver(text: &'static str) -> (String, Arc<AtomicBool>, thread::JoinHandle<()>) {
spawn_homeserver_nick(text, None)
}
fn spawn_homeserver_nick(
text: &'static str,
nick: Option<&'static str>,
) -> (String, Arc<AtomicBool>, thread::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
listener.set_nonblocking(true).expect("nonblocking");
let done = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&done);
let handle = thread::spawn(move || {
let deadline = Instant::now() + Duration::from_secs(8);
while !flag.load(Ordering::Relaxed) && Instant::now() < deadline {
match listener.accept() {
Ok((mut sock, _)) => {
let _ = sock.set_nonblocking(false);
let Some((headers, _)) = read_http(&mut sock) else {
continue;
};
let line = headers.lines().next().unwrap_or("");
let resp = if line.contains("/sync") && line.contains("since=") {
sync_empty()
} else if line.contains("/sync") {
sync_with_text_nick(text, nick)
} else if line.contains("/keys/") {
br#"{"one_time_key_counts":{"signed_curve25519":50}}"#.to_vec()
} else {
b"{}".to_vec()
};
write_http(&mut sock, "200 OK", &resp);
}
Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(15));
}
Err(_) => break,
}
}
});
(format!("http://{addr}"), done, handle)
}
fn spawn_routine() -> (
String,
Arc<Mutex<Vec<Hit>>>,
Arc<AtomicBool>,
thread::JoinHandle<()>,
) {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
listener.set_nonblocking(true).expect("nonblocking");
let hits = Arc::new(Mutex::new(Vec::new()));
let recorded = Arc::clone(&hits);
let done = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&done);
let handle = thread::spawn(move || {
let deadline = Instant::now() + Duration::from_secs(8);
while !flag.load(Ordering::Relaxed) && Instant::now() < deadline {
match listener.accept() {
Ok((mut sock, _)) => {
let _ = sock.set_nonblocking(false);
if let Some((headers, body)) = read_http(&mut sock) {
recorded.lock().expect("hits").push(Hit { headers, body });
let _ = sock.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
);
}
}
Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(15));
}
Err(_) => break,
}
}
});
(format!("http://{addr}/routine"), hits, done, handle)
}
fn stop(flag: &AtomicBool, handle: thread::JoinHandle<()>) {
flag.store(true, Ordering::Relaxed);
let _ = handle.join();
}
fn open_against(base: &str, dir: &Path, session: &str) -> OpenedStore {
let device = DeviceId::parse("DEVICE1").expect("device id");
OpenedStore::open(
dir,
session,
device,
"@alice:localhost",
"localhost",
plain(base, "fake-token"),
"fake-token",
)
.expect("open")
}
#[test]
fn inbound_room_text_goes_through_the_provider_chain_once() {
let (base, home_done, home) = spawn_homeserver("wake-chain");
let dir = temp_dir("wake-chain");
let inbox_dir = dir.0.join("inbox");
let mut store = open_against(&base, &dir.0, "session-c");
let session = provider::ProviderSession {
kind: provider::SessionKind::local(provider::ProviderKind::Codex),
session_id: "thread-c".into(),
nick: "builder".into(),
cwd: None,
headless: false,
};
let host = provider::HostEnv {
surface: provider::Surface::Local,
vendor: None,
os: "linux",
};
let config = provider::AdapterConfig {
inbox_dir: Some(inbox_dir.clone()),
spawn_program: Some("/bin/false".into()),
..provider::AdapterConfig::default()
};
let chain = provider::plan_chain(&session, &host, &config);
store.set_wake_chain(session, chain);
assert_eq!(store.wake_chain_ids()[0], "codex-app-server-turn");
store.drive(1_000, false).expect("drive");
store.drive(3_000, false).expect("drive again");
let route = store.take_wake_route();
let note = store.wake_note().unwrap_or("").to_string();
drop(store);
stop(&home_done, home);
assert_eq!(route.as_deref(), Some("inbox-queue queued"), "note={note}");
let letters = provider::inbox::drain(&inbox_dir);
assert_eq!(letters.len(), 1);
assert!(letters[0].letter.prompt.ends_with("wake-chain"));
assert!(letters[0].letter.prompt.contains("m4a-send --as builder"));
}
#[test]
fn inbound_room_text_hits_the_routine_once() {
let (base, home_done, home) = spawn_homeserver("wake-plain");
let (routine, hits, routine_done, routine_thread) = spawn_routine();
let dir = temp_dir("wake-once");
let mut store = open_against(&base, &dir.0, "session-a");
store.set_wake(SessionWake {
routine_url: Some(routine),
..SessionWake::default()
});
store.drive(1_000, false).expect("drive");
store.drive(3_000, false).expect("drive again");
let note = store.wake_note().unwrap_or("").to_string();
let saw_text = store.texts().iter().any(|text| text.body == "wake-plain");
let trace = store.http_trace();
drop(store);
stop(&routine_done, routine_thread);
stop(&home_done, home);
assert!(
saw_text,
"engine did not surface the inbound text; http={trace:?} note={note}"
);
let hits = hits.lock().expect("hits");
assert_eq!(
hits.len(),
1,
"routine was not hit exactly once; note={note}"
);
let hit = &hits[0];
assert!(
hit.headers.starts_with("POST /routine HTTP/1.1"),
"{}",
hit.headers.lines().next().unwrap_or("")
);
assert!(
!hit.headers.to_ascii_lowercase().contains("authorization"),
"no bearer was passed, so the routine post must not send one"
);
assert!(
!hit.headers
.to_ascii_lowercase()
.contains("x-automation-key"),
"no bearer was passed, so the routine post must not send a key"
);
assert!(!hit.headers.contains("/mail/send"));
assert!(!hit.headers.contains("/admin/listener"));
assert!(
hit.headers
.to_ascii_lowercase()
.contains("content-type: application/json"),
"{}",
hit.headers
);
assert_wake_json(
&hit.body,
"wake-plain",
"@bob:localhost",
"$m1:localhost",
None,
);
}
#[test]
fn routine_wake_stays_once_across_store_reopen() {
let (base, home_done, home) = spawn_homeserver("wake-persist");
let (routine, hits, routine_done, routine_thread) = spawn_routine();
let dir = temp_dir("wake-persist");
let mut store = open_against(&base, &dir.0, "session-a");
store.set_wake(SessionWake {
routine_url: Some(routine.clone()),
..SessionWake::default()
});
store.drive(1_000, false).expect("drive");
let ledger = dir.0.join("routine-woken");
assert!(
ledger.is_file(),
"successful wake must persist routine-woken"
);
drop(store);
let mut store = open_against(&base, &dir.0, "session-a");
store.set_wake(SessionWake {
routine_url: Some(routine),
..SessionWake::default()
});
store.drive(1_000, false).expect("drive again");
store.drive(3_000, false).expect("drive third");
let note = store.wake_note().unwrap_or("").to_string();
drop(store);
stop(&routine_done, routine_thread);
stop(&home_done, home);
let hits = hits.lock().expect("hits");
assert_eq!(
hits.len(),
1,
"reopen must not re-POST the same event; note={note} hits={}",
hits.len()
);
}
#[test]
fn inbound_room_text_posts_the_sender_nick_the_shell_already_has() {
let (base, home_done, home) = spawn_homeserver_nick("wake-named", Some("Alice"));
let (routine, hits, routine_done, routine_thread) = spawn_routine();
let dir = temp_dir("wake-nick");
let mut store = open_against(&base, &dir.0, "session-a");
store.set_wake(SessionWake {
routine_url: Some(routine),
..SessionWake::default()
});
store.drive(1_000, false).expect("drive");
let note = store.wake_note().unwrap_or("").to_string();
drop(store);
stop(&routine_done, routine_thread);
stop(&home_done, home);
let hits = hits.lock().expect("hits");
assert_eq!(
hits.len(),
1,
"routine was not hit exactly once; note={note}"
);
assert!(
!hits[0]
.headers
.to_ascii_lowercase()
.contains("authorization"),
"no bearer was passed"
);
assert_wake_json(
&hits[0].body,
"wake-named",
"@bob:localhost",
"$m1:localhost",
Some("Alice"),
);
}
#[test]
fn no_routine_url_posts_nothing() {
let (base, home_done, home) = spawn_homeserver("wake-plain");
let (_routine, hits, routine_done, routine_thread) = spawn_routine();
let dir = temp_dir("wake-none");
let mut store = open_against(&base, &dir.0, "session-a");
store.set_wake(SessionWake::default());
store.drive(1_000, false).expect("drive");
let saw_text = store.texts().iter().any(|text| text.body == "wake-plain");
let trace = store.http_trace();
drop(store);
stop(&routine_done, routine_thread);
stop(&home_done, home);
assert!(
saw_text,
"missing routine url dropped the inbound text; http={trace:?}"
);
let hits = hits.lock().expect("hits");
assert!(hits.is_empty(), "a post happened with no routine url");
}
#[test]
fn routine_bearer_header_is_sent_only_when_the_caller_passed_one() {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let server = thread::spawn(move || {
let (mut sock, _) = listener.accept().expect("accept");
let (headers, body) = read_http(&mut sock).expect("request");
let _ = sock
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n");
(headers, body)
});
let bearer = "env-bearer";
post_decrypted_with_bearer(
&format!("http://{addr}/routine"),
&DecryptedWake {
body: "letter",
from: "@bob:localhost",
nick: Some("Bob"),
event_id: "$m1:localhost",
..Default::default()
},
Some(bearer),
)
.expect("post");
let (headers, body) = server.join().expect("server");
let line = headers
.lines()
.find(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.expect("authorization header");
let (name, value) = line.split_once(':').expect("header");
assert!(name.eq_ignore_ascii_case("authorization"));
assert_eq!(value.trim(), format!("Bearer {bearer}"));
let automation = headers
.lines()
.find(|line| line.to_ascii_lowercase().starts_with("x-automation-key:"))
.expect("automation key header");
let (name, value) = automation.split_once(':').expect("header");
assert!(name.eq_ignore_ascii_case("x-automation-key"));
assert_eq!(value.trim(), bearer);
assert_eq!(
headers
.lines()
.filter(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.count(),
1
);
assert_eq!(
headers
.lines()
.filter(|line| line.to_ascii_lowercase().starts_with("x-automation-key:"))
.count(),
1
);
assert_wake_json(
&body,
"letter",
"@bob:localhost",
"$m1:localhost",
Some("Bob"),
);
}
#[test]
fn post_decrypted_omits_a_missing_or_blank_nick() {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let server = thread::spawn(move || {
let mut bodies = Vec::new();
for _ in 0..2 {
let (mut sock, _) = listener.accept().expect("accept");
let (headers, body) = read_http(&mut sock).expect("request");
assert!(
!headers.to_ascii_lowercase().contains("authorization"),
"no bearer was passed"
);
assert!(
!headers.to_ascii_lowercase().contains("x-automation-key"),
"no bearer was passed"
);
let _ = sock.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
);
bodies.push(body);
}
bodies
});
let url = format!("http://{addr}/routine");
post_decrypted(
&url,
&DecryptedWake {
body: "plain",
from: "@bob:localhost",
nick: None,
event_id: "$m1:localhost",
..Default::default()
},
)
.expect("missing nick");
post_decrypted(
&url,
&DecryptedWake {
body: "plain",
from: "@bob:localhost",
nick: Some(" "),
event_id: "$m1:localhost",
..Default::default()
},
)
.expect("blank nick");
let bodies = server.join().expect("server");
for body in &bodies {
assert_wake_json(body, "plain", "@bob:localhost", "$m1:localhost", None);
}
}
#[cfg(all(unix, feature = "wake-grok"))]
#[test]
fn inbound_room_text_prompts_the_leader_socket_once() {
use std::os::unix::net::{UnixListener, UnixStream};
fn frame_read(sock: &mut UnixStream) -> Option<Vec<u8>> {
let _ = sock.set_read_timeout(Some(Duration::from_secs(3)));
let mut len_buf = [0u8; 4];
sock.read_exact(&mut len_buf).ok()?;
let len = u32::from_be_bytes(len_buf) as usize;
if len > 1_000_000 {
return None;
}
let mut buf = vec![0u8; len];
sock.read_exact(&mut buf).ok()?;
Some(buf)
}
fn frame_write(sock: &mut UnixStream, value: &serde_json::Value) {
let bytes = serde_json::to_vec(value).expect("json");
let mut out = (bytes.len() as u32).to_be_bytes().to_vec();
out.extend_from_slice(&bytes);
sock.write_all(&out).expect("write");
sock.flush().expect("flush");
}
fn serve_one(sock: &mut UnixStream, prompts: &Mutex<Vec<String>>) {
let Some(bytes) = frame_read(sock) else {
return;
};
let register: serde_json::Value = serde_json::from_slice(&bytes).unwrap_or_default();
assert_eq!(register["type"], "register");
frame_write(
sock,
&serde_json::json!({"type": "registered", "ready": true}),
);
loop {
let Some(bytes) = frame_read(sock) else {
break;
};
let value: serde_json::Value = match serde_json::from_slice(&bytes) {
Ok(value) => value,
Err(_) => break,
};
if value.get("type").and_then(|item| item.as_str()) == Some("disconnect") {
break;
}
if value.get("type").and_then(|item| item.as_str()) != Some("acp") {
continue;
}
let payload = value
.get("payload")
.and_then(|item| item.as_str())
.unwrap_or("");
let inner: serde_json::Value = serde_json::from_str(payload).unwrap_or_default();
if inner.get("method").and_then(|item| item.as_str()) == Some("session/prompt") {
let text = inner["params"]["prompt"][0]["text"]
.as_str()
.unwrap_or("")
.to_string();
let session = inner["params"]["sessionId"]
.as_str()
.unwrap_or("")
.to_string();
prompts
.lock()
.expect("prompts")
.push(format!("{session} {text}"));
}
let id = inner.get("id").cloned().unwrap_or(serde_json::json!(null));
let body = serde_json::json!({"jsonrpc":"2.0","id": id, "result": {}}).to_string();
frame_write(sock, &serde_json::json!({"type":"acp","payload": body}));
}
}
let dir = temp_dir("leader");
let sock_dir = dir.0.join("sock");
std::fs::create_dir_all(&sock_dir).expect("dir");
let path = sock_dir.join("leader.sock");
let listener = UnixListener::bind(&path).expect("bind");
listener.set_nonblocking(true).expect("nonblocking");
let prompts = Arc::new(Mutex::new(Vec::<String>::new()));
let recorded = Arc::clone(&prompts);
let done = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&done);
let leader = thread::spawn(move || {
let deadline = Instant::now() + Duration::from_secs(8);
while !flag.load(Ordering::Relaxed) && Instant::now() < deadline {
match listener.accept() {
Ok((mut sock, _)) => {
let _ = sock.set_nonblocking(false);
serve_one(&mut sock, &recorded);
}
Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(15));
}
Err(_) => break,
}
}
});
let (base, home_done, home) = spawn_homeserver("wake-leader");
let mut store = open_against(&base, &dir.0.join("store"), "session-a");
store.set_wake(SessionWake {
leader_sock: Some(path),
leader_cwd: Some("/tmp".to_string()),
..SessionWake::default()
});
store.drive(1_000, false).expect("drive");
store.drive(3_000, false).expect("drive again");
let note = store.wake_note().unwrap_or("").to_string();
let saw_text = store.texts().iter().any(|text| text.body == "wake-leader");
drop(store);
stop(&done, leader);
stop(&home_done, home);
assert!(
saw_text,
"engine did not surface the inbound text; note={note}"
);
let prompts = prompts.lock().expect("prompts");
assert_eq!(
prompts.as_slice(),
["session-a wake-leader"],
"leader prompt was not the decrypted text once; note={note}"
);
}
}