wx-cli/.gsd/exec/a11d0ae5-8ad9-4047-bc7a-c49...

579 lines
19 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

=== 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 serverUnix 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 ===