mirror of https://github.com/jackwener/wx-cli.git
579 lines
19 KiB
Plaintext
579 lines
19 KiB
Plaintext
=== src/transport/mod.rs ===
|
||
//! Transport abstraction layer.
|
||
//!
|
||
//! Defines object-safe traits for listening/connecting over different
|
||
//! transport types (Unix socket, Windows named pipe, TCP) and a generic
|
||
//! connection handler that extracts the JSON-line protocol logic from
|
||
//! the platform-specific `handle_connection_unix/windows` in `server.rs`.
|
||
|
||
use std::future::Future;
|
||
use std::path::PathBuf;
|
||
use std::pin::Pin;
|
||
use std::net::SocketAddr;
|
||
use std::sync::Arc;
|
||
use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader};
|
||
use anyhow::Result;
|
||
|
||
use crate::daemon::cache::DbCache;
|
||
use crate::daemon::query::Names;
|
||
use crate::ipc::{Request, Response};
|
||
|
||
// ─── Transport address ───────────────────────────────────────────────────────
|
||
|
||
/// Unified transport address covering Unix socket, Windows named pipe, and TCP.
|
||
#[derive(Debug, Clone)]
|
||
pub enum TransportAddr {
|
||
Unix(PathBuf),
|
||
WindowsPipe(String),
|
||
Tcp(SocketAddr),
|
||
}
|
||
|
||
// ─── Traits ──────────────────────────────────────────────────────────────────
|
||
|
||
/// Object-safe trait for accepting incoming connections.
|
||
///
|
||
/// Each implementation provides its own concrete `Stream` type.
|
||
pub trait Listener {
|
||
type Stream: AsyncRead + AsyncWrite + Unpin + Send + 'static;
|
||
|
||
fn accept(&mut self) -> Pin<Box<dyn Future<Output = Result<Self::Stream>> + Send + '_>>;
|
||
}
|
||
|
||
/// Object-safe trait for initiating outgoing connections.
|
||
pub trait Connector {
|
||
type Stream: AsyncRead + AsyncWrite + Unpin + Send + 'static;
|
||
|
||
fn connect(
|
||
&self,
|
||
addr: &TransportAddr,
|
||
) -> Pin<Box<dyn Future<Output = Result<Self::Stream>> + Send + '_>>;
|
||
}
|
||
|
||
// ─── Generic connection handler ──────────────────────────────────────────────
|
||
|
||
/// Read one JSON line, parse as `Request`, dispatch, write one JSON-line `Response`.
|
||
///
|
||
/// Extracted from the duplicated `handle_connection_unix` / `handle_connection_windows`
|
||
/// in `server.rs`. The function is generic over the stream type so it works with
|
||
/// `UnixStream`, Windows named pipe stream, `TcpStream`, etc.
|
||
pub async fn handle_connection<S>(
|
||
mut stream: S,
|
||
db: &DbCache,
|
||
names: &Arc<tokio::sync::RwLock<Arc<Names>>>,
|
||
) -> Result<()>
|
||
where
|
||
S: AsyncRead + AsyncWrite + Unpin,
|
||
{
|
||
let (reader, mut writer) = tokio::io::split(&mut stream);
|
||
let mut lines = BufReader::new(reader).lines();
|
||
|
||
let line = match lines.next_line().await? {
|
||
Some(l) => l,
|
||
None => return Ok(()), // client closed without sending anything
|
||
};
|
||
|
||
// Parse request
|
||
let req: Request = match serde_json::from_str(&line) {
|
||
Ok(r) => r,
|
||
Err(e) => {
|
||
let resp = Response::err(format!("JSON 解析错误: {}", e));
|
||
writer.write_all(resp.to_json_line()?.as_bytes()).await?;
|
||
return Ok(());
|
||
}
|
||
};
|
||
|
||
let resp = dispatch(req, db, names).await;
|
||
writer.write_all(resp.to_json_line()?.as_bytes()).await?;
|
||
Ok(())
|
||
}
|
||
|
||
// ─── Dispatch (temporary copy from server.rs; will be shared in T02) ────────
|
||
|
||
async fn dispatch(
|
||
req: Request,
|
||
db: &DbCache,
|
||
names: &tokio::sync::RwLock<Arc<Names>>,
|
||
) -> Response {
|
||
use super::daemon::query;
|
||
|
||
let names_arc: Arc<Names> = {
|
||
let guard = names.read().await;
|
||
Arc::clone(&*guard)
|
||
};
|
||
|
||
match req {
|
||
Request::Ping => Response::ok(serde_json::json!({ "pong": true })),
|
||
Request::Sessions { limit } => {
|
||
match query::q_sessions(db, &names_arc, limit).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::History { chat, limit, offset, since, until, msg_type } => {
|
||
match query::q_history(db, &names_arc, &chat, limit, offset, since, until, msg_type).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::Search { keyword, chats, limit, since, until, msg_type } => {
|
||
match query::q_search(db, &names_arc, &keyword, chats, limit, since, until, msg_type).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::Contacts { query, limit } => {
|
||
match query::q_contacts(&names_arc, query.as_deref(), limit).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::Unread { limit, filter } => {
|
||
match query::q_unread(db, &names_arc, limit, filter).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::Members { chat } => {
|
||
match query::q_members(db, &names_arc, &chat).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::NewMessages { state, limit } => {
|
||
match query::q_new_messages(db, &names_arc, state, limit).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::Favorites { limit, fav_type, query } => {
|
||
match query::q_favorites(db, limit, fav_type, query).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::Stats { chat, since, until } => {
|
||
match query::q_stats(db, &names_arc, &chat, since, until).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::SnsNotifications { limit, since, until, include_read } => {
|
||
match query::q_sns_notifications(db, &names_arc, limit, since, until, include_read).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::SnsFeed { limit, since, until, user } => {
|
||
match query::q_sns_feed(db, &names_arc, limit, since, until, user.as_deref()).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
Request::SnsSearch { keyword, limit, since, until, user } => {
|
||
match query::q_sns_search(db, &names_arc, &keyword, limit, since, until, user.as_deref()).await {
|
||
Ok(v) => Response::ok(v),
|
||
Err(e) => Response::err(e.to_string()),
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
// ─── TCP implementations ────────────────────────────────────────────────────
|
||
|
||
/// TCP listener wrapping `tokio::net::TcpListener`.
|
||
pub struct TcpListener {
|
||
inner: tokio::net::TcpListener,
|
||
}
|
||
|
||
impl TcpListener {
|
||
pub async fn bind(addr: SocketAddr) -> Result<Self> {
|
||
let inner = tokio::net::TcpListener::bind(addr).await?;
|
||
Ok(Self { inner })
|
||
}
|
||
}
|
||
|
||
impl Listener for TcpListener {
|
||
type Stream = tokio::net::TcpStream;
|
||
|
||
fn accept(&mut self) -> Pin<Box<dyn Future<Output = Result<Self::Stream>> + Send + '_>> {
|
||
Box::pin(async {
|
||
let (stream, _addr) = self.inner.accept().await?;
|
||
Ok(stream)
|
||
})
|
||
}
|
||
}
|
||
|
||
/// TCP connector using `tokio::net::TcpStream`.
|
||
pub struct TcpConnector;
|
||
|
||
impl Connector for TcpConnector {
|
||
type Stream = tokio::net::TcpStream;
|
||
|
||
fn connect(
|
||
&self,
|
||
addr: &TransportAddr,
|
||
) -> Pin<Box<dyn Future<Output = Result<Self::Stream>> + Send + '_>> {
|
||
let addr = addr.clone();
|
||
Box::pin(async move {
|
||
match addr {
|
||
TransportAddr::Tcp(socket_addr) => {
|
||
let stream = tokio::net::TcpStream::connect(socket_addr).await?;
|
||
Ok(stream)
|
||
}
|
||
other => anyhow::bail!("TcpConnector 不支持 {:?},请使用对应的 Connector", other),
|
||
}
|
||
})
|
||
}
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
#[test]
|
||
fn transport_addr_variants() {
|
||
let unix = TransportAddr::Unix(PathBuf::from("/tmp/wx.sock"));
|
||
let tcp = TransportAddr::Tcp("127.0.0.1:8080".parse().unwrap());
|
||
let pipe = TransportAddr::WindowsPipe("wx-cli-daemon".to_string());
|
||
|
||
match unix {
|
||
TransportAddr::Unix(p) => assert_eq!(p, PathBuf::from("/tmp/wx.sock")),
|
||
_ => panic!("expected Unix"),
|
||
}
|
||
match tcp {
|
||
TransportAddr::Tcp(s) => assert_eq!(s.port(), 8080),
|
||
_ => panic!("expected Tcp"),
|
||
}
|
||
match pipe {
|
||
TransportAddr::WindowsPipe(s) => assert_eq!(s, "wx-cli-daemon"),
|
||
_ => panic!("expected WindowsPipe"),
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn tcp_connector_rejects_non_tcp_addr() {
|
||
// Verify at compile-time that TcpConnector implements Connector
|
||
fn assert_connector<T: Connector>() {}
|
||
assert_connector::<TcpConnector>();
|
||
}
|
||
|
||
#[test]
|
||
fn tcp_listener_implements_listener() {
|
||
fn assert_listener<T: Listener>() {}
|
||
assert_listener::<TcpListener>();
|
||
}
|
||
}
|
||
|
||
=== src/daemon/server.rs ===
|
||
use anyhow::Result;
|
||
use std::sync::Arc;
|
||
|
||
use crate::transport::{self, Listener};
|
||
use super::cache::DbCache;
|
||
use super::query::Names;
|
||
|
||
/// 启动 IPC server(Unix socket / Windows named pipe + 可选 TCP)
|
||
///
|
||
/// 当 `tcp_addr` 为 `Some` 时,同时监听 TCP 端口;daemon 在 local listener 退出时退出。
|
||
pub async fn serve(
|
||
db: Arc<DbCache>,
|
||
names: Arc<tokio::sync::RwLock<Arc<Names>>>,
|
||
tcp_addr: Option<&str>,
|
||
) -> Result<()> {
|
||
// TCP 先启动为后台任务
|
||
if let Some(addr) = tcp_addr {
|
||
let socket_addr: std::net::SocketAddr = addr.parse().map_err(|e| {
|
||
anyhow::anyhow!("TCP 地址解析失败 '{}': {}", addr, e)
|
||
})?;
|
||
let db_tcp = Arc::clone(&db);
|
||
let names_tcp = Arc::clone(&names);
|
||
tokio::spawn(async move {
|
||
if let Err(e) = serve_tcp(socket_addr, db_tcp, names_tcp).await {
|
||
eprintln!("[server] TCP 监听错误: {}", e);
|
||
}
|
||
});
|
||
}
|
||
|
||
#[cfg(unix)]
|
||
serve_unix(db, names).await?;
|
||
#[cfg(windows)]
|
||
serve_windows(db, names).await?;
|
||
Ok(())
|
||
}
|
||
|
||
async fn serve_tcp(
|
||
addr: std::net::SocketAddr,
|
||
db: Arc<DbCache>,
|
||
names: Arc<tokio::sync::RwLock<Arc<Names>>>,
|
||
) -> Result<()> {
|
||
let listener = transport::TcpListener::bind(addr).await?;
|
||
eprintln!("[server] 监听 TCP {}", addr);
|
||
|
||
// TcpListener::accept 返回 Pin<Box<dyn Future>>,需要 Box::pin 包装循环
|
||
let mut listener = listener;
|
||
loop {
|
||
let stream = listener.accept().await?;
|
||
let db2 = Arc::clone(&db);
|
||
let names2 = Arc::clone(&names);
|
||
tokio::spawn(async move {
|
||
if let Err(e) = transport::handle_connection(stream, &db2, &names2).await {
|
||
eprintln!("[server] 连接处理错误: {}", e);
|
||
}
|
||
});
|
||
}
|
||
}
|
||
|
||
#[cfg(unix)]
|
||
async fn serve_unix(
|
||
db: Arc<DbCache>,
|
||
names: Arc<tokio::sync::RwLock<Arc<Names>>>,
|
||
) -> Result<()> {
|
||
use tokio::net::UnixListener;
|
||
let sock_path = crate::config::sock_path();
|
||
|
||
// 删除旧 socket 文件
|
||
if sock_path.exists() {
|
||
let _ = tokio::fs::remove_file(&sock_path).await;
|
||
}
|
||
|
||
let listener = UnixListener::bind(&sock_path)?;
|
||
// 设置权限 0600
|
||
#[cfg(unix)]
|
||
{
|
||
use std::os::unix::fs::PermissionsExt;
|
||
std::fs::set_permissions(&sock_path, std::fs::Permissions::from_mode(0o600))?;
|
||
}
|
||
|
||
eprintln!("[server] 监听 {}", sock_path.display());
|
||
|
||
loop {
|
||
let (stream, _) = listener.accept().await?;
|
||
let db2 = Arc::clone(&db);
|
||
let names2 = Arc::clone(&names);
|
||
|
||
tokio::spawn(async move {
|
||
if let Err(e) = transport::handle_connection(stream, &db2, &names2).await {
|
||
eprintln!("[server] 连接处理错误: {}", e);
|
||
}
|
||
});
|
||
}
|
||
}
|
||
|
||
#[cfg(windows)]
|
||
async fn serve_windows(
|
||
db: Arc<DbCache>,
|
||
names: Arc<tokio::sync::RwLock<Arc<Names>>>,
|
||
) -> Result<()> {
|
||
use interprocess::local_socket::{
|
||
tokio::prelude::*, GenericNamespaced, ListenerOptions,
|
||
};
|
||
|
||
// interprocess 的 GenericNamespaced 在 Windows 上会自动拼接 `\\.\pipe\` 前缀,
|
||
// 这里必须传相对名;client 端用 `\\.\pipe\wx-cli-daemon` 直接打开可以对上
|
||
let name = "wx-cli-daemon".to_ns_name::<GenericNamespaced>()?;
|
||
let opts = ListenerOptions::new().name(name);
|
||
let listener = opts.create_tokio()?;
|
||
|
||
eprintln!("[server] 监听 \\\\.\\pipe\\wx-cli-daemon");
|
||
|
||
loop {
|
||
let conn = listener.accept().await?;
|
||
let db2 = Arc::clone(&db);
|
||
let names2 = Arc::clone(&names);
|
||
|
||
tokio::spawn(async move {
|
||
if let Err(e) = transport::handle_connection(conn, &db2, &names2).await {
|
||
eprintln!("[server] 连接处理错误: {}", e);
|
||
}
|
||
});
|
||
}
|
||
}
|
||
|
||
=== src/daemon/mod.rs ===
|
||
pub mod cache;
|
||
pub mod query;
|
||
pub mod server;
|
||
|
||
use anyhow::Result;
|
||
use std::collections::HashMap;
|
||
use std::sync::Arc;
|
||
|
||
use crate::config;
|
||
|
||
/// daemon 入口
|
||
///
|
||
/// 当 WX_DAEMON_MODE 环境变量设置时,main() 调用此函数
|
||
pub fn run() {
|
||
let rt = tokio::runtime::Runtime::new().expect("无法创建 tokio runtime");
|
||
if let Err(e) = rt.block_on(start_daemon(None)) {
|
||
eprintln!("[daemon] 启动失败: {}", e);
|
||
std::process::exit(1);
|
||
}
|
||
}
|
||
|
||
/// 从 CLI `wx daemon start [--tcp ADDR]` 调用
|
||
///
|
||
/// 查找当前可执行文件路径,设置 WX_DAEMON_MODE=1,后台启动新进程。
|
||
pub fn run_start(tcp_addr: Option<String>) -> Result<()> {
|
||
let exe = std::env::current_exe()?;
|
||
let log = config::log_path();
|
||
|
||
let mut cmd = std::process::Command::new(&exe);
|
||
cmd.env("WX_DAEMON_MODE", "1");
|
||
if let Some(addr) = &tcp_addr {
|
||
cmd.env("WX_DAEMON_TCP_ADDR", addr);
|
||
}
|
||
// 日志重定向
|
||
let log_file = std::fs::OpenOptions::new()
|
||
.create(true)
|
||
.append(true)
|
||
.open(&log)?;
|
||
cmd.stdout(log_file.try_clone()?).stderr(log_file);
|
||
|
||
#[cfg(unix)]
|
||
{
|
||
use std::os::unix::process::CommandExt;
|
||
unsafe { cmd.pre_exec(|| {
|
||
libc::setsid();
|
||
Ok(())
|
||
}) };
|
||
}
|
||
|
||
let child = cmd.spawn()?;
|
||
let pid = child.id();
|
||
eprintln!("[daemon] 已启动 daemon 进程 (PID {})", pid);
|
||
Ok(())
|
||
}
|
||
|
||
/// daemon 核心启动逻辑(被 run() 和 WX_DAEMON_MODE 路径共享)
|
||
pub async fn start_daemon(tcp_addr: Option<String>) -> Result<()> {
|
||
// 确保工作目录存在
|
||
let cli_dir = config::cli_dir();
|
||
tokio::fs::create_dir_all(&cli_dir).await?;
|
||
tokio::fs::create_dir_all(config::cache_dir()).await?;
|
||
|
||
// 写 PID 文件
|
||
let pid = std::process::id();
|
||
tokio::fs::write(config::pid_path(), pid.to_string()).await?;
|
||
|
||
// 注册 SIGTERM / SIGINT 处理
|
||
setup_signal_handler().await;
|
||
|
||
eprintln!("[daemon] wx-daemon 启动 (PID {})", pid);
|
||
|
||
// 加载配置
|
||
let cfg = config::load_config()?;
|
||
eprintln!("[daemon] DB_DIR: {}", cfg.db_dir.display());
|
||
|
||
// 加载密钥
|
||
let keys_content = tokio::fs::read_to_string(&cfg.keys_file).await
|
||
.map_err(|e| anyhow::anyhow!("读取密钥文件 {:?} 失败: {}", cfg.keys_file, e))?;
|
||
let keys_raw: serde_json::Value = serde_json::from_str(&keys_content)?;
|
||
let all_keys = extract_keys(&keys_raw);
|
||
eprintln!("[daemon] 密钥数量: {}", all_keys.len());
|
||
|
||
// 初始化 DbCache
|
||
let db = Arc::new(cache::DbCache::new(cfg.db_dir.clone(), all_keys.clone()).await?);
|
||
|
||
// 收集消息 DB 列表
|
||
let msg_db_keys: Vec<String> = all_keys.keys()
|
||
.filter(|k| {
|
||
let k = k.replace('\\', "/");
|
||
k.contains("message/message_") && k.ends_with(".db")
|
||
&& !k.contains("_fts") && !k.contains("_resource")
|
||
})
|
||
.cloned()
|
||
.collect();
|
||
|
||
// 预热:加载联系人 + 解密 session.db
|
||
eprintln!("[daemon] 预热...");
|
||
let names_raw = query::load_names(&*db).await.unwrap_or_else(|e| {
|
||
eprintln!("[daemon] 加载联系人失败: {}", e);
|
||
query::Names {
|
||
map: HashMap::new(),
|
||
md5_to_uname: HashMap::new(),
|
||
msg_db_keys: Vec::new(),
|
||
verify_flags: HashMap::new(),
|
||
}
|
||
});
|
||
let mut names = names_raw;
|
||
names.msg_db_keys = msg_db_keys;
|
||
|
||
let _ = db.get("session/session.db").await;
|
||
let _ = db.get("sns/sns.db").await;
|
||
eprintln!("[daemon] 预热完成,联系人 {} 个", names.map.len());
|
||
|
||
// 包一层内部 Arc
|
||
let names_arc = Arc::new(tokio::sync::RwLock::new(Arc::new(names)));
|
||
|
||
// 检查环境变量中的 TCP 地址(WX_DAEMON_MODE 路径下通过 env 传入)
|
||
let effective_tcp_addr = tcp_addr.or_else(|| std::env::var("WX_DAEMON_TCP_ADDR").ok());
|
||
|
||
// 启动 IPC server(阻塞)
|
||
server::serve(Arc::clone(&db), Arc::clone(&names_arc), effective_tcp_addr.as_deref()).await?;
|
||
|
||
// 正常退出时清理(signal 路径下由 cleanup_and_exit 处理,不会走到这里)
|
||
#[allow(unreachable_code)]
|
||
{
|
||
let _ = std::fs::remove_file(config::sock_path());
|
||
let _ = std::fs::remove_file(config::pid_path());
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
|
||
/// 从 all_keys.json 提取 rel_key -> enc_key 映射
|
||
///
|
||
/// 兼容两种格式:
|
||
/// - `{ "rel/path.db": { "enc_key": "hex" } }`(Python 版原生格式)
|
||
/// - `{ "rel/path.db": "hex" }`(简化格式)
|
||
fn extract_keys(json: &serde_json::Value) -> HashMap<String, String> {
|
||
let mut result = HashMap::new();
|
||
if let Some(obj) = json.as_object() {
|
||
for (k, v) in obj {
|
||
if k.starts_with('_') { continue; }
|
||
let enc_key = if let Some(s) = v.as_str() {
|
||
s.to_string()
|
||
} else if let Some(obj2) = v.as_object() {
|
||
obj2.get("enc_key")
|
||
.and_then(|e| e.as_str())
|
||
.unwrap_or_default()
|
||
.to_string()
|
||
} else {
|
||
continue;
|
||
};
|
||
if !enc_key.is_empty() {
|
||
// 统一路径分隔符
|
||
let rel = k.replace('\\', "/");
|
||
result.insert(rel, enc_key);
|
||
}
|
||
}
|
||
}
|
||
result
|
||
}
|
||
|
||
/// 设置信号处理(Unix: SIGTERM/SIGINT)
|
||
async fn setup_signal_handler() {
|
||
#[cfg(unix)]
|
||
tokio::spawn(async move {
|
||
use tokio::signal::unix::{signal, SignalKind};
|
||
let mut term = signal(SignalKind::terminate()).expect("无法监听 SIGTERM");
|
||
let mut int = signal(SignalKind::interrupt()).expect("无法监听 SIGINT");
|
||
tokio::select! {
|
||
_ = term.recv() => {},
|
||
_ = int.recv() => {},
|
||
}
|
||
cleanup_and_exit();
|
||
});
|
||
}
|
||
|
||
#[cfg(unix)]
|
||
fn cleanup_and_exit() {
|
||
// 仅清理 local socket 文件,TCP 端口由 OS 自动回收
|
||
let _ = std::fs::remove_file(config::sock_path());
|
||
let _ = std::fs::remove_file(config::pid_path());
|
||
std::process::exit(0);
|
||
}
|
||
|
||
=== src/daemon/cli.rs ===
|