PicoBot/src/client/mod.rs

437 lines
15 KiB
Rust

pub use crate::protocol::{WsInbound, WsOutbound, serialize_inbound, serialize_outbound};
mod oneshot;
mod tui;
pub use oneshot::{RunOptions, read_run_prompt, run_once};
use crate::client::tui::app::{App, MessageRole};
use crate::client::tui::event::{
handle_key_event, handle_paste, request_history, request_session_list, send,
};
use crate::client::tui::ui::render_ui;
use crossterm::{
event::{self, DisableBracketedPaste, EnableBracketedPaste, Event, KeyEventKind},
execute,
terminal::{EnterAlternateScreen, LeaveAlternateScreen, disable_raw_mode, enable_raw_mode},
};
use futures_util::StreamExt;
use ratatui::{Terminal, prelude::CrosstermBackend};
use std::io;
use std::{fs, path::PathBuf};
use tokio_tungstenite::{
connect_async,
tungstenite::{Message, client::IntoClientRequest, http::header},
};
pub async fn run(
gateway_url: &str,
pair_code: Option<&str>,
) -> Result<(), Box<dyn std::error::Error>> {
let client_id = load_or_create_client_id();
let separator = if gateway_url.contains('?') { '&' } else { '?' };
let connect_url = format!("{gateway_url}{separator}client_id={client_id}");
let token = if let Some(code) = pair_code {
let token = exchange_pairing_code(gateway_url, code).await?;
save_auth_token(&token)?;
Some(token)
} else {
load_auth_token()
};
let mut request = connect_url.into_client_request()?;
if let Some(token) = &token {
request
.headers_mut()
.insert(header::AUTHORIZATION, format!("Bearer {token}").parse()?);
}
let (ws_stream, _) = connect_async(request).await.map_err(|error| {
format!(
"gateway connection failed: {error}. If pairing is required, run `picobot pair` then `picobot chat --pair-code <CODE>`"
)
})?;
tracing::info!("Connected to gateway");
let (ws_sender, ws_receiver) = ws_stream.split();
let mut app = App::new();
app.http_base_url = gateway_http_base_url(gateway_url)?;
app.auth_token = token;
app.client_id = client_id;
app.ws_sender = Some(ws_sender);
app.ws_receiver = Some(ws_receiver);
enable_raw_mode()?;
let mut stdout = io::stdout();
execute!(stdout, EnterAlternateScreen, EnableBracketedPaste)?;
let backend = CrosstermBackend::new(stdout);
let mut terminal = Terminal::new(backend)?;
terminal.clear()?;
let result = run_app(&mut terminal, app).await;
// Cleanup terminal, ignore errors
let _ = execute!(
terminal.backend_mut(),
DisableBracketedPaste,
LeaveAlternateScreen
);
let _ = disable_raw_mode();
let _ = terminal.show_cursor();
result
}
fn gateway_http_base_url(gateway_url: &str) -> Result<String, Box<dyn std::error::Error>> {
let mut url = reqwest::Url::parse(gateway_url)?;
let scheme = match url.scheme() {
"ws" => "http",
"wss" => "https",
"http" => "http",
"https" => "https",
other => return Err(format!("unsupported gateway URL scheme: {other}").into()),
};
url.set_scheme(scheme)
.map_err(|_| "failed to set gateway URL scheme")?;
url.set_path("");
url.set_query(None);
url.set_fragment(None);
Ok(url.to_string().trim_end_matches('/').to_string())
}
pub async fn reload_gateway(gateway_url: &str) -> Result<String, Box<dyn std::error::Error>> {
let base = gateway_http_base_url(gateway_url)?;
let mut request = reqwest::Client::new().post(format!("{base}/api/config/reload"));
if let Some(token) = load_auth_token() {
request = request.bearer_auth(token);
}
let response = request.send().await?;
let status = response.status();
let body: serde_json::Value = response.json().await?;
if !status.is_success() {
return Err(body
.get("error")
.and_then(serde_json::Value::as_str)
.unwrap_or("configuration reload failed")
.to_string()
.into());
}
Ok(body
.get("message")
.and_then(serde_json::Value::as_str)
.unwrap_or("configuration reload scheduled")
.to_string())
}
async fn exchange_pairing_code(
gateway_url: &str,
code: &str,
) -> Result<String, Box<dyn std::error::Error>> {
let mut url = reqwest::Url::parse(gateway_url)?;
let scheme = match url.scheme() {
"ws" => "http",
"wss" => "https",
"http" => "http",
"https" => "https",
other => return Err(format!("unsupported gateway URL scheme: {other}").into()),
};
url.set_scheme(scheme)
.map_err(|_| "failed to set gateway URL scheme")?;
url.set_path("/api/auth/pair");
url.set_query(None);
url.set_fragment(None);
let response = reqwest::Client::new()
.post(url)
.json(&serde_json::json!({ "code": code }))
.send()
.await?;
let status = response.status();
let body: serde_json::Value = response.json().await?;
if !status.is_success() {
return Err(body
.get("error")
.and_then(serde_json::Value::as_str)
.unwrap_or("pairing failed")
.to_string()
.into());
}
body.get("token")
.and_then(serde_json::Value::as_str)
.map(str::to_string)
.ok_or_else(|| "gateway did not return an auth token".into())
}
fn auth_token_path() -> Option<PathBuf> {
dirs::home_dir().map(|home| home.join(".picobot").join("tui_auth_token"))
}
fn load_auth_token() -> Option<String> {
fs::read_to_string(auth_token_path()?)
.ok()
.map(|token| token.trim().to_string())
.filter(|token| !token.is_empty())
}
fn save_auth_token(token: &str) -> Result<(), Box<dyn std::error::Error>> {
let path = auth_token_path().ok_or("home directory is unavailable")?;
let parent = path.parent().ok_or("invalid auth token path")?;
fs::create_dir_all(parent)?;
#[cfg(unix)]
{
use std::io::Write;
use std::os::unix::fs::{OpenOptionsExt, PermissionsExt};
let mut file = fs::OpenOptions::new()
.create(true)
.truncate(true)
.write(true)
.mode(0o600)
.open(&path)?;
file.write_all(token.as_bytes())?;
file.sync_all()?;
fs::set_permissions(path, fs::Permissions::from_mode(0o600))?;
}
#[cfg(not(unix))]
fs::write(&path, token)?;
Ok(())
}
fn load_or_create_client_id() -> String {
let generated = uuid::Uuid::new_v4().simple().to_string();
let Some(home) = dirs::home_dir() else {
return generated;
};
let dir = home.join(".picobot");
let path: PathBuf = dir.join("tui_client_id");
if let Ok(value) = fs::read_to_string(&path) {
let value = value.trim();
if !value.is_empty()
&& value.len() <= 64
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-' || byte == b'_')
{
return value.to_string();
}
}
if fs::create_dir_all(dir).is_ok() {
let _ = fs::write(path, &generated);
}
generated
}
async fn run_app(
terminal: &mut Terminal<CrosstermBackend<io::Stdout>>,
mut app: App,
) -> Result<(), Box<dyn std::error::Error>> {
let mut ws_receiver = app.ws_receiver.take().unwrap();
let mut event_reader = event::EventStream::new();
let mut ws_open = true;
send(&mut app, WsInbound::GetSlashCommands).await;
request_session_list(&mut app).await;
loop {
terminal.draw(|f| render_ui(f, &app))?;
tokio::select! {
msg = ws_receiver.next(), if ws_open => {
match msg {
Some(Ok(Message::Text(text))) => {
if let Ok(outbound) = serde_json::from_str::<WsOutbound>(&text) {
handle_ws_message(&mut app, outbound).await;
}
}
Some(Ok(Message::Close(_))) | None => {
tracing::info!("Gateway disconnected");
app.connected = false;
app.ws_sender = None;
ws_open = false;
app.status_message = Some("Gateway 连接已关闭;按两次 Ctrl+C 退出".to_string());
}
Some(Err(error)) => {
app.connected = false;
app.ws_sender = None;
ws_open = false;
app.status_message = Some(format!("Gateway 连接错误:{error}"));
}
_ => {}
}
}
event_result = event_reader.next() => {
match event_result {
Some(Ok(Event::Key(key))) if key.kind != KeyEventKind::Release => {
handle_key_event(&mut app, key).await;
}
Some(Ok(Event::Paste(text))) => handle_paste(&mut app, &text),
Some(Err(error)) => {
app.status_message = Some(format!("终端输入错误:{error}"));
}
_ => {}
}
}
}
if app.should_quit {
break;
}
}
Ok(())
}
async fn handle_ws_message(app: &mut App, outbound: WsOutbound) {
match outbound {
WsOutbound::TurnUpdated { snapshot } => {
let terminal = snapshot.status != crate::session::TurnStatus::Running;
let completed = snapshot.status == crate::session::TurnStatus::Completed;
let session_id = snapshot.session_id.clone();
if terminal {
app.pending_responses = app.pending_responses.saturating_sub(1);
}
if app.apply_turn_snapshot(snapshot) {
if terminal {
app.status_message = None;
if app.current_session_id.as_deref() == Some(&session_id) {
if !completed {
request_history(app, session_id).await;
}
} else {
app.status_message = Some("另一个会话已完成响应".to_string());
}
request_session_list(app).await;
} else {
app.status_message = Some("正在生成回复…".to_string());
}
} else if terminal && app.current_session_id.as_deref() != Some(&session_id) {
app.status_message = Some("另一个会话已完成响应".to_string());
}
}
WsOutbound::TurnCommitted {
session_id,
history_revision,
messages,
} => {
app.apply_turn_commit(&session_id, history_revision, messages);
}
WsOutbound::AssistantResponse {
id,
content,
attachments,
session_id,
..
} => {
app.pending_responses = app.pending_responses.saturating_sub(1);
app.status_message = None;
if session_id
.as_ref()
.is_none_or(|session_id| app.current_session_id.as_ref() == Some(session_id))
{
app.add_message_with_attachments(id, MessageRole::Assistant, content, attachments);
if let Some(current) = app.current_session_id.clone() {
request_history(app, current).await;
}
} else {
app.status_message = Some("另一个会话已完成响应".to_string());
}
request_session_list(app).await;
}
WsOutbound::Error { message, .. } => {
app.pending_responses = app.pending_responses.saturating_sub(1);
app.status_message = Some(message.clone());
app.add_message(MessageRole::System, format!("Error: {}", message));
}
WsOutbound::SessionEstablished {
session_id,
capabilities,
} => {
app.connected = true;
app.file_transfer_supported =
capabilities.iter().any(|value| value == "file_transfer_v1");
app.set_current_session(Some(session_id.clone()));
request_history(app, session_id).await;
request_session_list(app).await;
}
WsOutbound::SessionCreated { session_id, .. } => {
app.set_current_session(Some(session_id.clone()));
app.status_message = None;
request_history(app, session_id).await;
request_session_list(app).await;
}
WsOutbound::SessionList {
sessions,
current_session_id,
} => {
app.set_sessions(sessions);
if let Some(id) = current_session_id {
let changed = app.current_session_id.as_deref() != Some(&id);
app.set_current_session(Some(id.clone()));
if changed {
request_history(app, id).await;
}
}
}
WsOutbound::SessionLoaded { session_id, .. } => {
app.set_current_session(Some(session_id.clone()));
request_history(app, session_id).await;
request_session_list(app).await;
}
WsOutbound::SessionHistory {
session_id,
messages,
} => app.set_history(&session_id, messages),
// The first Todo UI is WebUI-only. CLI keeps receiving ordinary task
// notifications and may inspect plans through /todo.
WsOutbound::SessionPlan { .. } | WsOutbound::PlanUpdated { .. } => {}
WsOutbound::SessionRenamed { session_id, title } => {
if let Some(session) = app
.sessions
.iter_mut()
.find(|session| session.session_id == session_id)
{
session.title = title;
}
request_session_list(app).await;
}
WsOutbound::SessionArchived { session_id } => {
app.sessions
.retain(|session| session.session_id != session_id);
request_session_list(app).await;
}
WsOutbound::SessionDeleted { session_id } => {
if app.current_session_id.as_ref() == Some(&session_id) {
app.set_current_session(None);
}
app.sessions
.retain(|session| session.session_id != session_id);
request_session_list(app).await;
}
WsOutbound::HistoryCleared { session_id } => {
if app.current_session_id.as_deref() == Some(&session_id) {
app.messages.clear();
}
app.status_message = None;
request_session_list(app).await;
}
WsOutbound::SlashCommandsList { commands } => {
app.set_commands(commands);
}
WsOutbound::Pong => {}
WsOutbound::CommandExecuted { message } => {
app.pending_responses = app.pending_responses.saturating_sub(1);
app.status_message = None;
app.add_message(MessageRole::System, message);
request_session_list(app).await;
}
WsOutbound::SystemNotification {
content,
session_id,
} => {
if session_id
.as_ref()
.is_none_or(|session_id| app.current_session_id.as_ref() == Some(session_id))
{
app.add_message(MessageRole::System, content);
}
}
}
}