增加oneshot客户端,解决一次性对话和测试的支持。

This commit is contained in:
xiaoxixi 2026-07-19 23:18:39 +08:00
parent 711b07af83
commit 762fd16e3f
8 changed files with 630 additions and 20 deletions

View File

@ -7,6 +7,7 @@ This file is the operational contract for coding agents working in this reposito
- `cargo build` — build the binary
- `cargo run -- gateway` — start gateway server (binds `127.0.0.1:19876` by default)
- `cargo run -- chat` — connect to gateway as CLI client (default `ws://127.0.0.1:19876/ws`)
- `cargo run -- run "prompt"` — send one prompt through Gateway, print the terminal Turn, and exit; stdin, JSON, verbose progress, and timeout modes are available
- `docker compose up -d` — start the container with Gateway bound/published on `0.0.0.0:19876`; override `PICOBOT_GATEWAY_HOST`, `PICOBOT_PUBLISH_HOST`, or `PICOBOT_GATEWAY_PORT` as needed
- WebUI — start Gateway, then open `http://127.0.0.1:19876/`; no separate frontend build is required
- `cd webui && npm ci && npm run check && npm run build` — validate the Svelte WebUI independently (Node.js 20+); its local `dist/` is ignored
@ -19,6 +20,7 @@ This file is the operational contract for coding agents working in this reposito
- `.env` files use a custom parser, not dotenv: load `<config-dir>/.env`, then `<workspace_dir>/.env`, while pre-existing process variables remain highest priority; config placeholders `<VAR_NAME>` use the merged values
- Config example: `resources/templates/config.example.json` (released to `~/.picobot/` on first run)
- CLI TUI identity is stored in `~/.picobot/tui_client_id`; it is a non-secret stable chat scope used to restore dialogs across reconnects
- One-shot `run` uses a unique chat scope per invocation; for loopback Gateway URLs it authenticates `/ws` with `~/.picobot/web_admin_token`, while remote URLs use the existing paired CLI token
## Tests
@ -39,6 +41,7 @@ This file is the operational contract for coding agents working in this reposito
- **Gateway mode** (`cargo run -- gateway`): HTTP/WebSocket server; owns `GatewayState` which holds all services
- **Client mode** (`cargo run -- chat`): TUI chat client; connects to gateway via WebSocket, purely for user interaction
- **One-shot client mode** (`cargo run -- run "prompt"`): isolated CLI chat scope; connects to Gateway, waits for a terminal Turn, prints it, and exits
### Core Data Flow
@ -91,7 +94,7 @@ Scheduler → SessionManager scheduled execution → AgentLoop → Scheduler del
- **WebUI slash completion** must consume the existing `get_slash_commands` WebSocket response; do not duplicate the backend command list in frontend source
- **WebUI chat rendering** sanitizes Markdown before inserting HTML; durable `turn_committed` deltas calibrate normal terminal Turns without a full history reload, and history must preserve structured tool-call metadata so calls and results remain independently collapsible
- **WebUI/TUI file transfer** streams bytes over authenticated HTTP and sends only short-lived upload IDs/attachment metadata over WebSocket; messages persist local media paths without guaranteeing later availability, and client responses must never expose those paths
- **WebUI authentication** protects every management API and `/ws`; only static pairing assets, public health/status, and pairing submission may bypass device auth. Pair-code issuance requires both a real loopback peer and the filesystem-held admin token; never put bearer tokens in URLs or logs
- **WebUI authentication** protects every management API and `/ws`; only static pairing assets, public health/status, and pairing submission may bypass device auth. Pair-code issuance requires both a real loopback peer and the filesystem-held admin token. The same conjunction may authenticate only `/ws` for local one-shot `run`; it must never authorize management APIs. Never put bearer or admin tokens in URLs or logs
- **Providers** are pure HTTP clients; no bus/session/channel awareness
- **Provider reasoning state** is private replay data: persist it, replay it only to the matching provider, and never expose it to clients, channels, or logs
- **Tools** are executed by `AgentLoop`; they receive raw arguments and normally return text. Tools that produce model-consumable media use the structured `execute_with_media` side channel; model capability checks and provider content-block serialization stay outside tools
@ -136,3 +139,6 @@ Scheduler → SessionManager scheduled execution → AgentLoop → Scheduler del
- `docs/ARCHITECTURE.md` — maintainer-facing runtime design, invariants, lifecycle, and extension guidance
- `AGENTS.md` — concise operational rules for repository agents
- `resources/skills/about-picobot/references/` — runtime knowledge shipped to PicoBot; update it only when the assistant's built-in product knowledge must change
## Version Management
- 在每次功能变化、架构变化后,适当地更新整个产品的版本号

View File

@ -9,6 +9,7 @@ PicoBot 是一个用 Rust 编写的个人 AI 助手运行时。它在本地启
## 适合做什么
- 在终端里和本地 AI 助手持续对话。
- 从脚本或命令行发送一条任务,等待完整的模型/工具循环后只输出最终结果。
- 在 TUI 或浏览器中实时查看正文、思考过程和工具执行状态,并在完成后收敛到持久化历史。
- 在浏览器中查看日志、任务和记忆,修改运行配置与助手档案。
- 复杂任务可创建 session 级 Todo 计划,把不同子项并行委托给多个子 Agent聊天页侧栏实时显示进度。
@ -124,9 +125,22 @@ docker compose up -d
cargo run -- chat
```
CLI 默认连接 `ws://127.0.0.1:19876/ws`。首次使用先运行 `picobot pair`,再执行 `picobot chat --pair-code <CODE>`;客户端令牌会以 `0600` 权限保存到 `~/.picobot/tui_auth_token`。如需指定地址,可使用 `--gateway-url`
CLI 默认连接 `ws://127.0.0.1:19876/ws`TUI 首次使用先运行 `picobot pair`,再执行 `picobot chat --pair-code <CODE>`;客户端令牌会以 `0600` 权限保存到 `~/.picobot/tui_auth_token`。如需指定地址,可使用 `--gateway-url`
### 5.1 使用 WebUI
### 5.1 一次性执行
`run` 通过 Gateway 发送一条消息,复用正常的 SessionManager、AgentLoop 和工具调用流程,收到 Turn 终态后打印最终回复并退出:
```bash
picobot run "检查这个项目并总结测试结果"
printf '使用浏览器打开 example.com 并返回页面标题\n' | picobot run
```
默认情况下 stdout 只包含最终回复,便于管道和脚本消费。`--verbose` 把阶段和工具状态写到 stderr`--json` 输出包含 session、turn、状态、正文、usage 和错误的一行 JSON`--timeout` 设置最大等待秒数。超时或按下 Ctrl-C 时,客户端会先向当前会话发送 `/stop`
连接本机回环地址时不需要人工配对:`run` 自动读取 `~/.picobot/web_admin_token`Gateway 只有在真实 TCP 对端也是回环地址时才允许该凭据访问 `/ws`。每次调用使用独立的临时 chat scope不会替换正在运行的 TUI 连接。连接远程 Gateway 时仍使用 `~/.picobot/tui_auth_token` 中已有的配对令牌。
### 5.2 使用 WebUI
Gateway 启动后直接打开:
@ -163,7 +177,7 @@ WebUI 随二进制嵌入,不需要 Node.js、npm 或单独部署静态文件
配置接口会掩码 API Key、secret、password 和 token保留 `********` 再保存不会覆盖原密钥。运行配置采用原子写入并在 Gateway 重启后生效,`USER.md``AGENTS.md` 则会用于后续构建的 Agent 上下文。
WebUI 默认启用设备配对鉴权,管理 API 与 `/ws` 都拒绝未配对客户端;静态配对页、公开健康检查和配对提交接口除外。令牌只以 SHA-256 哈希写入 `~/.picobot/web_auth.json`,本地配对码管理密钥位于权限为 `0600``~/.picobot/web_admin_token`。鉴权不提供传输加密;如果通过 `--host 0.0.0.0`、反向代理或端口转发暴露 Gateway仍必须使用 TLS。可通过 `gateway.require_pairing=false` 显式关闭配对,但不建议在非隔离环境使用。
WebUI 默认启用设备配对鉴权,管理 API 与 `/ws` 都拒绝未配对客户端;静态配对页、公开健康检查和配对提交接口除外。唯一的 WebSocket 例外是本机 `picobot run`:请求必须同时来自真实回环对端并持有权限为 `0600``~/.picobot/web_admin_token`,该管理令牌不能绕过任何管理 API 的设备鉴权。配对令牌只以 SHA-256 哈希写入 `~/.picobot/web_auth.json`。鉴权不提供传输加密;如果通过 `--host 0.0.0.0`、反向代理或端口转发暴露 Gateway仍必须使用 TLS。可通过 `gateway.require_pairing=false` 显式关闭配对,但不建议在非隔离环境使用。
#### WebUI 开发

33
config.json Normal file
View File

@ -0,0 +1,33 @@
{
"providers": {
"aliyun": {
"type": "openai",
"base_url": "https://example.invalid/v1",
"api_key": "test-only-not-a-real-key",
"extra_headers": {}
}
},
"models": {
"qwen-plus": {
"model_id": "qwen-plus",
"temperature": 0.0,
"max_tokens": 100,
"input_type": ["text"]
}
},
"agents": {
"default": {
"provider": "aliyun",
"model": "qwen-plus",
"max_tool_iterations": 20,
"token_limit": 128000
}
},
"gateway": {
"host": "127.0.0.1",
"port": 19876,
"require_pairing": true
},
"channels": {},
"workspace_dir": "/tmp/picobot-test-workspace"
}

View File

@ -18,12 +18,13 @@ PicoBot 是一个单进程、异步、可扩展的个人 AI 助手运行时。
## 2. 运行模式与进程边界
PicoBot 只有一个二进制,提供两种模式:
PicoBot 只有一个二进制,提供三种运行模式:
| 模式 | 入口 | 职责 |
|------|------|------|
| Gateway | `cargo run -- gateway` | 组装服务、监听 HTTP/WebSocket、提供嵌入式 WebUI运行渠道、会话、调度器和后台任务 |
| CLI client | `cargo run -- chat` | 运行 Ratatui UI通过 WebSocket 使用 Gateway不持有业务状态 |
| One-shot client | `cargo run -- run "prompt"` | 使用独立临时 chat scope 通过 WebSocket 提交一条消息,等待 Turn 终态,输出结果后退出 |
Linux 上可通过 `picobot service install/start/stop/status/restart/uninstall` 管理 systemd 用户服务。unit 固定为 `picobot.service`,其主进程仍是普通 Gateway 模式,不引入额外 daemon/fork 层;异常退出由 systemd 按 `Restart=on-failure` 拉起。
@ -31,6 +32,8 @@ Linux 上可通过 `picobot service install/start/stop/status/restart/uninstall`
CLI TUI 在 `~/.picobot/tui_client_id` 保存非敏感客户端标识,并通过 WebSocket 查询参数 `client_id` 传给 Gateway。`cli_chat` 以该标识作为稳定 chat scope重连时恢复内存中的当前 dialogGateway 重启后则恢复该 scope 最近活跃的未归档 dialog。无效或缺失的标识会退化为连接级随机 scope。
One-shot client 不绕过 Gateway 直接调用 Provider。每次 `run` 生成独立的 `run-<uuid>` scope通过相同的 `cli_chat`、MessageBus、SessionManager、AgentLoop 和 Turn delivery 路径执行;它不复用 TUI scope因而不会替换同一 scope 的活动 WebSocket。默认 stdout 只投影终态 Assistant blocks进度写到 stderr超时或 Ctrl-C 会先在当前 scope 发送 `/stop`
Gateway 启动时先从配置目录 `.env`、workspace `.env` 和既有进程环境合并启动变量,再初始化日志并切换进程工作目录到 `workspace_dir`。优先级为进程环境 > workspace `.env` > 配置目录 `.env`;配置目录层先用于定位 workspaceworkspace 层不得重定向自身位置。环境文件只在单线程启动阶段写入进程环境,不能移到后台任务启动之后。切换完成后所有相对路径都应按 workspace 解释,不能假设仍位于源码仓库。
## 3. 组件关系
@ -244,7 +247,7 @@ WebUI/TUI 文件字节通过受鉴权的 HTTP 接口流式传输WebSocket 只
工具默认通过 `ToolResult` 返回文本;需要把图片等产物交给模型时,通过 `Tool::execute_with_media` 返回文本和结构化 `MediaRef`。工具只负责经过自身路径策略校验后声明媒体,不感知当前模型或 Provider。`AgentLoop` 仅将最新连续工具结果批次的媒体交给 `MediaHandlerRegistry`,旧工具媒体只回放文本和路径,避免历史 Base64 膨胀。OpenAI-compatible Provider 保持 `tool` 结果为文本,并在完整工具批次后构造仅存在于请求内的临时多模态 `user` 消息Anthropic Provider 将媒体放入对应 `tool_result.content`。媒体加载、格式或能力检查失败必须降级成文本,不得使历史记录不可读取。
`AuthManager` 默认保护所有管理 API 和 `/ws`。静态资源、`/health``/api/auth/status``/api/auth/pair` 保持公开,使未配对浏览器只能加载配对界面。`picobot pair` 使用权限为 `0600` 的本机管理密钥调用仅接受真实回环连接的 `/api/auth/code`;反向代理即使从回环连接也无法在没有该密钥时签发代码。配对码为 8 位、5 分钟有效、单次消费,并按来源实施失败锁定。浏览器收到 HttpOnly、SameSite=Strict CookieCLI 使用 Bearer token服务端仅持久化 SHA-256 哈希。`--revoke-all` 的持久化成功后才清空内存令牌,活动 WebSocket 每 5 秒复核身份并回收已撤销连接。
`AuthManager` 默认保护所有管理 API 和 `/ws`。静态资源、`/health``/api/auth/status``/api/auth/pair` 保持公开,使未配对浏览器只能加载配对界面。`picobot pair` 使用权限为 `0600` 的本机管理密钥调用仅接受真实回环连接的 `/api/auth/code`;反向代理即使从回环连接也无法在没有该密钥时签发代码。本机 `picobot run` 可用同一个管理密钥直接认证 `/ws`,但中间件必须同时验证请求路径严格等于 `/ws``ConnectInfo` 中的真实 TCP 对端为回环地址;这一身份不能访问管理 API。远程 `run` 与 TUI 一样使用已配对的 Bearer token。配对码为 8 位、5 分钟有效、单次消费,并按来源实施失败锁定。浏览器收到 HttpOnly、SameSite=Strict CookieCLI 使用 Bearer token服务端仅持久化 SHA-256 哈希。`--revoke-all` 的持久化成功后才清空内存令牌,活动 WebSocket 每 5 秒复核身份并回收已撤销连接。
同源 `/api/*` 管理接口只提供显式白名单能力:

View File

@ -1,7 +1,10 @@
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,

428
src/client/oneshot.rs Normal file
View File

@ -0,0 +1,428 @@
use super::{WsInbound, WsOutbound, load_auth_token};
use crate::config::get_user_config_dir;
use crate::gateway::auth::ADMIN_TOKEN_HEADER;
use crate::session::{ToolStatus, TurnBlock, TurnPhase, TurnSnapshot, TurnStatus};
use futures_util::{SinkExt, StreamExt};
use serde::Serialize;
use std::collections::HashMap;
use std::io::{self, IsTerminal, Read, Write};
use std::net::IpAddr;
use std::time::Duration;
use tokio_tungstenite::connect_async;
use tokio_tungstenite::tungstenite::{
Message,
client::IntoClientRequest,
http::{HeaderValue, header},
};
const MAX_RUN_PROMPT_BYTES: usize = 1024 * 1024;
type DynError = Box<dyn std::error::Error>;
#[derive(Debug, Clone, Copy)]
pub struct RunOptions {
pub timeout: Duration,
pub json: bool,
pub verbose: bool,
}
#[derive(Debug, Serialize)]
struct RunOutput {
session_id: String,
turn_id: String,
status: TurnStatus,
content: String,
usage: Option<crate::providers::Usage>,
error: Option<String>,
}
pub fn read_run_prompt(parts: Vec<String>) -> Result<String, DynError> {
if !parts.is_empty() {
return validate_prompt(parts.join(" "));
}
if io::stdin().is_terminal() {
return Err("provide a prompt as arguments or pipe it on stdin".into());
}
let stdin = io::stdin();
let mut locked = stdin.lock();
read_prompt_from(&mut locked)
}
pub async fn run_once(
gateway_url: &str,
prompt: String,
options: RunOptions,
) -> Result<(), Box<dyn std::error::Error>> {
if options.timeout.is_zero() {
return Err("run timeout must be greater than zero".into());
}
let prompt = validate_prompt(prompt)?;
let client_id = format!("run-{}", uuid::Uuid::new_v4().simple());
let (connect_url, local_gateway) = websocket_url(gateway_url, &client_id)?;
let admin_token = local_gateway
.then(|| std::fs::read_to_string(get_user_config_dir().join("web_admin_token")).ok())
.flatten()
.map(|token| token.trim().to_string())
.filter(|token| !token.is_empty());
let bearer_token = admin_token.is_none().then(load_auth_token).flatten();
let mut request = connect_url.into_client_request()?;
if let Some(token) = &admin_token {
let mut value = HeaderValue::from_str(token)?;
value.set_sensitive(true);
request.headers_mut().insert(ADMIN_TOKEN_HEADER, value);
} else if let Some(token) = &bearer_token {
let mut value = HeaderValue::from_str(&format!("Bearer {token}"))?;
value.set_sensitive(true);
request.headers_mut().insert(header::AUTHORIZATION, value);
}
let (stream, _) = connect_async(request).await.map_err(|error| {
if local_gateway && admin_token.is_none() {
format!(
"gateway connection failed: {error}; local admin token is unavailable at {}",
get_user_config_dir().join("web_admin_token").display()
)
} else {
format!(
"gateway connection failed: {error}. Remote gateways require an existing paired CLI token"
)
}
})?;
let (mut sender, mut receiver) = stream.split();
let operation = async {
let session_id = loop {
match receiver.next().await {
Some(Ok(Message::Text(text))) => match serde_json::from_str::<WsOutbound>(&text)? {
WsOutbound::SessionEstablished { session_id, .. } => break session_id,
WsOutbound::Error { code, message } => {
return Err::<RunOutput, DynError>(
format!("gateway error {code}: {message}").into(),
);
}
_ => {}
},
Some(Ok(Message::Close(_))) | None => {
return Err("gateway closed before establishing a session".into());
}
Some(Err(error)) => return Err(error.into()),
_ => {}
}
};
let input = WsInbound::UserInput {
content: prompt,
upload_ids: Vec::new(),
channel: None,
chat_id: None,
sender_id: None,
};
sender
.send(Message::Text(serde_json::to_string(&input)?.into()))
.await?;
let mut turn_id = None;
let mut last_phase = None;
let mut tool_states: HashMap<String, (String, ToolStatus)> = HashMap::new();
loop {
match receiver.next().await {
Some(Ok(Message::Text(text))) => match serde_json::from_str::<WsOutbound>(&text)? {
WsOutbound::TurnUpdated { snapshot }
if snapshot.session_id == session_id
&& turn_id.as_ref().is_none_or(|id| id == &snapshot.id.0) =>
{
turn_id.get_or_insert_with(|| snapshot.id.0.clone());
if options.verbose {
report_progress(&snapshot, &mut last_phase, &mut tool_states);
}
if snapshot.status != TurnStatus::Running {
break Ok(output_from_snapshot(snapshot));
}
}
WsOutbound::Error { code, message } => {
break Err(format!("gateway error {code}: {message}").into());
}
_ => {}
},
Some(Ok(Message::Close(_))) | None => {
break Err("gateway closed before the run completed".into());
}
Some(Err(error)) => break Err(error.into()),
_ => {}
}
}
};
let output = tokio::select! {
result = tokio::time::timeout(options.timeout, operation) => {
match result {
Ok(result) => result?,
Err(_) => {
send_stop(&mut sender).await;
return Err(format!("run timed out after {} seconds", options.timeout.as_secs()).into());
}
}
}
signal = tokio::signal::ctrl_c() => {
send_stop(&mut sender).await;
signal?;
return Err("run cancelled".into());
}
};
render_output(&output, options.json)?;
match output.status {
TurnStatus::Completed => Ok(()),
TurnStatus::Cancelled => Err(output
.error
.unwrap_or_else(|| "run cancelled".to_string())
.into()),
TurnStatus::Failed => Err(output
.error
.unwrap_or_else(|| "run failed".to_string())
.into()),
TurnStatus::Running => Err("gateway returned a non-terminal run result".into()),
}
}
async fn send_stop<S>(sender: &mut S)
where
S: futures_util::Sink<Message> + Unpin,
{
let stop = WsInbound::UserInput {
content: "/stop".to_string(),
upload_ids: Vec::new(),
channel: None,
chat_id: None,
sender_id: None,
};
if let Ok(text) = serde_json::to_string(&stop) {
let _ = sender.send(Message::Text(text.into())).await;
let _ = sender.flush().await;
}
}
fn websocket_url(
gateway_url: &str,
client_id: &str,
) -> Result<(String, bool), Box<dyn std::error::Error>> {
let mut url = reqwest::Url::parse(gateway_url)?;
let scheme = match url.scheme() {
"ws" => "ws",
"wss" => "wss",
"http" => "ws",
"https" => "wss",
other => return Err(format!("unsupported gateway URL scheme: {other}").into()),
};
url.set_scheme(scheme)
.map_err(|_| "failed to set gateway URL scheme")?;
if url.path().is_empty() || url.path() == "/" {
url.set_path("/ws");
}
url.query_pairs_mut().append_pair("client_id", client_id);
let local_gateway = url.host_str().is_some_and(|host| {
let host = host
.strip_prefix('[')
.and_then(|value| value.strip_suffix(']'))
.unwrap_or(host);
host.eq_ignore_ascii_case("localhost")
|| host
.parse::<IpAddr>()
.is_ok_and(|address| address.is_loopback())
});
Ok((url.to_string(), local_gateway))
}
fn read_prompt_from(reader: &mut impl Read) -> Result<String, Box<dyn std::error::Error>> {
let mut bytes = Vec::new();
reader
.take((MAX_RUN_PROMPT_BYTES + 1) as u64)
.read_to_end(&mut bytes)?;
if bytes.len() > MAX_RUN_PROMPT_BYTES {
return Err(format!("prompt exceeds {MAX_RUN_PROMPT_BYTES} bytes").into());
}
let prompt = String::from_utf8(bytes)?;
validate_prompt(prompt.trim_end_matches(['\r', '\n']).to_string())
}
fn validate_prompt(prompt: String) -> Result<String, Box<dyn std::error::Error>> {
if prompt.len() > MAX_RUN_PROMPT_BYTES {
return Err(format!("prompt exceeds {MAX_RUN_PROMPT_BYTES} bytes").into());
}
if prompt.trim().is_empty() {
return Err("prompt is empty".into());
}
Ok(prompt)
}
fn output_from_snapshot(snapshot: TurnSnapshot) -> RunOutput {
RunOutput {
session_id: snapshot.session_id,
turn_id: snapshot.id.0,
status: snapshot.status,
content: assistant_text(&snapshot.blocks),
usage: snapshot.usage,
error: snapshot.error,
}
}
fn assistant_text(blocks: &[TurnBlock]) -> String {
blocks
.iter()
.filter_map(|block| match block {
TurnBlock::Assistant { text, .. } if !text.is_empty() => Some(text.as_str()),
_ => None,
})
.collect::<Vec<_>>()
.join("\n\n")
}
fn report_progress(
snapshot: &TurnSnapshot,
last_phase: &mut Option<TurnPhase>,
tool_states: &mut HashMap<String, (String, ToolStatus)>,
) {
if last_phase.as_ref() != Some(&snapshot.phase) {
eprintln!("[phase: {}]", phase_name(snapshot.phase));
*last_phase = Some(snapshot.phase);
}
for block in &snapshot.blocks {
let TurnBlock::Tool {
id, name, status, ..
} = block
else {
continue;
};
let current = (name.clone(), *status);
if tool_states.get(id) != Some(&current) {
eprintln!("[tool: {name}: {}]", tool_status_name(*status));
tool_states.insert(id.clone(), current);
}
}
}
fn phase_name(phase: TurnPhase) -> &'static str {
match phase {
TurnPhase::Queued => "queued",
TurnPhase::Reasoning => "reasoning",
TurnPhase::Responding => "responding",
TurnPhase::Acting => "acting",
TurnPhase::Finalizing => "finalizing",
}
}
fn tool_status_name(status: ToolStatus) -> &'static str {
match status {
ToolStatus::Running => "running",
ToolStatus::Completed => "completed",
ToolStatus::Failed => "failed",
}
}
fn render_output(output: &RunOutput, json: bool) -> Result<(), Box<dyn std::error::Error>> {
if json {
println!("{}", serde_json::to_string(output)?);
} else if output.status == TurnStatus::Completed {
print!("{}", output.content);
if !output.content.ends_with('\n') {
println!();
}
io::stdout().flush()?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::session::{BlockId, TurnId};
#[test]
fn positional_prompt_parts_are_joined() {
assert_eq!(
validate_prompt(["hello", "world"].join(" ")).unwrap(),
"hello world"
);
}
#[test]
fn stdin_prompt_preserves_lines_and_trims_terminal_newline() {
let mut input = "first\nsecond\n".as_bytes();
assert_eq!(read_prompt_from(&mut input).unwrap(), "first\nsecond");
}
#[test]
fn empty_and_oversized_prompts_are_rejected() {
assert!(validate_prompt(" \n".to_string()).is_err());
assert!(validate_prompt("x".repeat(MAX_RUN_PROMPT_BYTES + 1)).is_err());
}
#[test]
fn assistant_blocks_are_joined_without_reasoning_or_tools() {
let blocks = vec![
TurnBlock::Reasoning {
id: BlockId("reasoning".to_string()),
iteration: 0,
text: "hidden".to_string(),
},
TurnBlock::Assistant {
id: BlockId("answer-1".to_string()),
iteration: 0,
text: "hello".to_string(),
},
TurnBlock::Tool {
id: "tool".to_string(),
iteration: 0,
name: "bash".to_string(),
arguments: serde_json::json!({}),
status: ToolStatus::Completed,
preview: None,
},
TurnBlock::Assistant {
id: BlockId("answer-2".to_string()),
iteration: 1,
text: "world".to_string(),
},
];
assert_eq!(assistant_text(&blocks), "hello\n\nworld");
}
#[test]
fn loopback_gate_uses_the_url_host() {
assert!(
websocket_url("ws://127.0.0.1:19876/ws", "run-id")
.unwrap()
.1
);
assert!(websocket_url("ws://[::1]:19876/ws", "run-id").unwrap().1);
assert!(websocket_url("http://localhost:19876", "run-id").unwrap().1);
assert!(
!websocket_url("wss://gateway.example/ws", "run-id")
.unwrap()
.1
);
}
#[test]
fn terminal_snapshot_becomes_script_output() {
let output = output_from_snapshot(TurnSnapshot {
id: TurnId("turn".to_string()),
session_id: "session".to_string(),
message_id: "message".to_string(),
revision: 1,
status: TurnStatus::Completed,
phase: TurnPhase::Finalizing,
blocks: vec![TurnBlock::Assistant {
id: BlockId("answer".to_string()),
iteration: 0,
text: "done".to_string(),
}],
usage: None,
error: None,
});
assert_eq!(output.turn_id, "turn");
assert_eq!(output.content, "done");
assert_eq!(output.status, TurnStatus::Completed);
}
}

View File

@ -21,6 +21,7 @@ const MAX_FAILED_ATTEMPTS: u32 = 5;
const MAX_TRACKED_CLIENTS: usize = 4096;
const MAX_PAIRED_TOKENS: usize = 128;
const AUTH_COOKIE: &str = "picobot_auth";
pub const ADMIN_TOKEN_HEADER: &str = "X-Picobot-Admin-Token";
#[derive(Debug, Clone, Serialize, Deserialize)]
struct AuthStore {
@ -65,8 +66,12 @@ pub struct AuthManager {
state: Arc<Mutex<AuthState>>,
}
#[derive(Debug, Clone)]
pub struct AuthIdentity(pub Option<String>);
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AuthIdentity {
PairingDisabled,
Paired { token_hash: String },
LocalAdmin,
}
#[derive(Debug)]
pub enum PairError {
@ -110,7 +115,7 @@ impl AuthManager {
pub async fn authenticate(&self, token: Option<&str>) -> Option<AuthIdentity> {
if !self.required {
return Some(AuthIdentity(None));
return Some(AuthIdentity::PairingDisabled);
}
let hash = hash_token(token?);
self.state
@ -118,17 +123,17 @@ impl AuthManager {
.await
.token_hashes
.contains(&hash)
.then_some(AuthIdentity(Some(hash)))
.then_some(AuthIdentity::Paired { token_hash: hash })
}
pub async fn identity_is_active(&self, identity: &AuthIdentity) -> bool {
if !self.required {
return true;
match identity {
AuthIdentity::PairingDisabled => !self.required,
AuthIdentity::LocalAdmin => true,
AuthIdentity::Paired { token_hash } => {
self.required && self.state.lock().await.token_hashes.contains(token_hash)
}
}
let Some(hash) = identity.0.as_ref() else {
return false;
};
self.state.lock().await.token_hashes.contains(hash)
}
pub fn authenticate_admin(&self, token: Option<&str>) -> bool {
@ -210,9 +215,24 @@ pub async fn require_auth(
mut request: Request<Body>,
next: Next,
) -> Response {
let identity = auth
let mut identity = auth
.authenticate(token_from_headers(request.headers()))
.await;
if identity.is_none()
&& request.uri().path() == "/ws"
&& request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.is_some_and(|ConnectInfo(peer)| peer.ip().is_loopback())
&& auth.authenticate_admin(
request
.headers()
.get(ADMIN_TOKEN_HEADER)
.and_then(|value| value.to_str().ok()),
)
{
identity = Some(AuthIdentity::LocalAdmin);
}
let Some(identity) = identity else {
return (
StatusCode::UNAUTHORIZED,
@ -314,7 +334,7 @@ pub async fn issue_code(
headers: HeaderMap,
) -> Response {
let admin_token = headers
.get("X-Picobot-Admin-Token")
.get(ADMIN_TOKEN_HEADER)
.and_then(|value| value.to_str().ok());
if !peer.ip().is_loopback() || !state.auth.authenticate_admin(admin_token) {
return (
@ -488,7 +508,7 @@ async fn load_or_create_admin_token(path: &Path) -> Result<String, std::io::Erro
#[cfg(test)]
mod tests {
use super::*;
use axum::{Router, middleware, routing};
use axum::{Extension, Router, middleware, routing};
use tower::ServiceExt;
#[tokio::test]
@ -593,4 +613,61 @@ mod tests {
.unwrap();
assert_eq!(authorized.status(), StatusCode::OK);
}
#[tokio::test]
async fn local_admin_auth_is_limited_to_loopback_websockets() {
let dir = tempfile::tempdir().unwrap();
let manager = AuthManager::load(true, dir.path().join("auth.json"))
.await
.unwrap();
let admin_token = tokio::fs::read_to_string(dir.path().join("web_admin_token"))
.await
.unwrap();
let app = Router::new()
.route(
"/ws",
routing::get(|Extension(identity): Extension<AuthIdentity>| async move {
assert_eq!(identity, AuthIdentity::LocalAdmin);
StatusCode::OK
}),
)
.route("/protected", routing::get(|| async { StatusCode::OK }))
.route_layer(middleware::from_fn_with_state(manager, require_auth));
let mut local_ws = Request::get("/ws")
.header(ADMIN_TOKEN_HEADER, admin_token.trim())
.body(Body::empty())
.unwrap();
local_ws
.extensions_mut()
.insert(ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 42000))));
assert_eq!(
app.clone().oneshot(local_ws).await.unwrap().status(),
StatusCode::OK
);
let mut remote_ws = Request::get("/ws")
.header(ADMIN_TOKEN_HEADER, admin_token.trim())
.body(Body::empty())
.unwrap();
remote_ws
.extensions_mut()
.insert(ConnectInfo(SocketAddr::from(([192, 0, 2, 10], 42000))));
assert_eq!(
app.clone().oneshot(remote_ws).await.unwrap().status(),
StatusCode::UNAUTHORIZED
);
let mut local_api = Request::get("/protected")
.header(ADMIN_TOKEN_HEADER, admin_token.trim())
.body(Body::empty())
.unwrap();
local_api
.extensions_mut()
.insert(ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 42000))));
assert_eq!(
app.oneshot(local_api).await.unwrap().status(),
StatusCode::UNAUTHORIZED
);
}
}

View File

@ -30,6 +30,23 @@ enum Command {
#[arg(long)]
pair_code: Option<String>,
},
/// Send one prompt through the gateway, print the final response, and exit
Run {
/// Prompt text; when omitted, read it from stdin
prompt: Vec<String>,
/// Gateway WebSocket or HTTP URL
#[arg(long)]
gateway_url: Option<String>,
/// Maximum time to wait for the turn, in seconds
#[arg(long, default_value_t = 300)]
timeout: u64,
/// Print the terminal turn as one JSON object
#[arg(long)]
json: bool,
/// Print phase and tool progress to stderr
#[arg(long)]
verbose: bool,
},
/// Start gateway server
Gateway {
/// Host to bind to
@ -77,6 +94,32 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
.unwrap_or_else(|| "ws://127.0.0.1:19876/ws".to_string());
picobot::client::run(&url, pair_code.as_deref()).await?;
}
Command::Run {
prompt,
gateway_url,
timeout,
json,
verbose,
} => {
if timeout == 0 {
return Err("--timeout must be greater than zero".into());
}
let config = picobot::config::Config::load_default().ok();
let url = gateway_url
.or_else(|| config.as_ref().map(|c| c.client.gateway_url.clone()))
.unwrap_or_else(|| "ws://127.0.0.1:19876/ws".to_string());
let prompt = picobot::client::read_run_prompt(prompt)?;
picobot::client::run_once(
&url,
prompt,
picobot::client::RunOptions {
timeout: std::time::Duration::from_secs(timeout),
json,
verbose,
},
)
.await?;
}
Command::Gateway { host, port } => {
picobot::gateway::run(host, port).await?;
}
@ -101,7 +144,10 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
})?;
let response = reqwest::Client::new()
.post(endpoint)
.header("X-Picobot-Admin-Token", admin_token.trim())
.header(
picobot::gateway::auth::ADMIN_TOKEN_HEADER,
admin_token.trim(),
)
.send()
.await?;
let status = response.status();