xagent-pi 0.2.3

Self-contained local brain (chat UI + API + SSE) for the Pi agent, tunneled into xagent-service.
//! Tunnel client: connects OUT to xagent-service `/tunnel` and proxies each
//! forwarded HTTP request to the local brain HTTP server.
//!
//! Each request is handled CONCURRENTLY (spawned), so a long-lived SSE stream
//! on `/api/events` does not block subsequent requests. Frames to the shared
//! WebSocket sink are serialized per-frame via a mutex.

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) -> Result<()> {
    let client = reqwest::Client::builder().build()?;
    let url = tunnel_url(&service_url, &account)?;
    loop {
        match one_session(&url, client.clone(), &brain_origin).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) -> 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;
        };
        // Handle concurrently so a long-lived SSE request can't starve others.
        let sink = Arc::clone(&sink);
        let origin = origin.to_string();
        let client = client.clone();
        tokio::spawn(async move {
            if let Err(e) = serve(&sink, &client, &origin, 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,
    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);
    }
    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(())
}