use anyhow::Result;
use crate::config::AccountConfig;
use base64::{engine::general_purpose::STANDARD as B64, Engine as _};
use futures_util::{stream::SplitSink, SinkExt, StreamExt};
use xagent_service::frame::Frame;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::net::TcpStream;
use tokio::sync::Mutex;
use tokio_tungstenite::{
connect_async, MaybeTlsStream, WebSocketStream,
};
use tokio_tungstenite::tungstenite::Message;
type WsSink = SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>;
pub async fn run(service_url: String, account: AccountConfig, brain_origin: String, tunnel_token: String) -> Result<()> {
let client = reqwest::Client::builder().build()?;
let url = tunnel_url(&service_url, &account)?;
loop {
match one_session(&url, client.clone(), &brain_origin, &tunnel_token).await {
Ok(()) => log::info!("tunnel session ended cleanly"),
Err(e) => log::warn!("tunnel session error: {e}"),
}
log::info!("reconnecting in 2s");
tokio::time::sleep(Duration::from_secs(2)).await;
}
}
pub async fn register_account(service_url: &str, account: &AccountConfig) -> Result<()> {
let client = reqwest::Client::builder().build()?;
let base = service_url
.replace("wss://", "https://")
.replace("ws://", "http://");
let response = client
.post(format!("{}/register", base.trim_end_matches('/')))
.json(&serde_json::json!({ "username": account.username, "password": account.password }))
.send()
.await?;
match response.status().as_u16() {
200 | 201 => Ok(()),
409 => anyhow::bail!("用户名已存在: {}", account.username),
status => anyhow::bail!("注册失败,服务返回 {status}"),
}
}
fn tunnel_url(service_url: &str, account: &AccountConfig) -> Result<String> {
let mut url = reqwest::Url::parse(service_url)?;
let scheme = match url.scheme() {
"http" | "ws" => "ws",
"https" | "wss" => "wss",
scheme => anyhow::bail!("不支持的服务协议: {scheme}"),
};
url.set_scheme(scheme).map_err(|_| anyhow::anyhow!("无法设置隧道协议"))?;
url.set_path("/tunnel");
url.set_query(None);
url.query_pairs_mut()
.append_pair("username", &account.username)
.append_pair("password", &account.password);
Ok(url.into())
}
async fn one_session(url: &str, client: reqwest::Client, origin: &str, tunnel_token: &str) -> Result<()> {
log::info!("connecting tunnel: {url}");
let (ws, _resp) = connect_async(url).await?;
log::info!("tunnel connected → {origin}");
let (sink, mut stream) = ws.split();
let sink: Arc<Mutex<WsSink>> = Arc::new(Mutex::new(sink));
while let Some(msg) = stream.next().await {
let req_frame = match msg {
Ok(Message::Text(t)) => match serde_json::from_str::<Frame>(t.as_str()) {
Ok(f) => f,
Err(_) => continue,
},
Ok(Message::Close(_)) | Err(_) => break,
_ => continue,
};
let Frame::Req {
id,
method,
path,
headers,
body_b64,
} = req_frame
else {
continue;
};
let sink = Arc::clone(&sink);
let origin = origin.to_string();
let client = client.clone();
let tunnel_token = tunnel_token.to_string();
tokio::spawn(async move {
if let Err(e) = serve(&sink, &client, &origin, &tunnel_token, id, method, path, headers, body_b64).await {
log::warn!("serve error: {e}");
}
});
}
Ok(())
}
async fn send_frame(sink: &Arc<Mutex<WsSink>>, f: &Frame) -> Result<()> {
let text = serde_json::to_string(f)?;
sink.lock()
.await
.send(Message::Text(text))
.await
.map_err(|e| anyhow::anyhow!("ws send: {e}"))
}
#[allow(clippy::too_many_arguments)]
async fn serve(
sink: &Arc<Mutex<WsSink>>,
client: &reqwest::Client,
origin: &str,
tunnel_token: &str,
id: String,
method: String,
path: String,
headers: HashMap<String, String>,
body_b64: Option<String>,
) -> Result<()> {
let suffix = if path.starts_with('/') { path.clone() } else { format!("/{path}") };
let url = format!("{}{}", origin.trim_end_matches('/'), suffix);
let m = method
.parse::<reqwest::Method>()
.unwrap_or(reqwest::Method::GET);
let mut rb = client.request(m, &url);
for (k, v) in &headers {
rb = rb.header(k, v);
}
rb = rb.header("x-xagent-tunnel-auth", tunnel_token);
if let Some(b64) = body_b64 {
if let Ok(bytes) = B64.decode(&b64) {
rb = rb.body(bytes);
}
}
let resp = match rb.send().await {
Ok(r) => r,
Err(e) => {
send_frame(sink, &Frame::Res { id: id.clone(), status: 502, headers: HashMap::new() }).await?;
let _ = send_frame(sink, &Frame::StreamEnd { id: id.clone() }).await;
log::warn!("brain fetch failed: {e}");
return Ok(());
}
};
let status = resp.status().as_u16();
let mut h = HashMap::new();
for (k, v) in resp.headers().iter() {
if let Ok(val) = v.to_str() {
h.insert(k.as_str().to_string(), val.to_string());
}
}
send_frame(sink, &Frame::Res { id: id.clone(), status, headers: h }).await?;
let mut stream = resp.bytes_stream();
while let Some(chunk) = stream.next().await {
match chunk {
Ok(bytes) if !bytes.is_empty() => {
send_frame(
sink,
&Frame::StreamData { id: id.clone(), body_b64: B64.encode(&bytes) },
)
.await?;
}
Ok(_) => {}
Err(e) => {
log::warn!("brain stream error: {e}");
break;
}
}
}
send_frame(sink, &Frame::StreamEnd { id }).await?;
Ok(())
}