Ouvre le point d'intégration du serveur embarqué desktop : `run_embedded_with_core` reçoit un `Arc<BackendCore>` déjà construit, pour que l'adapter HTTP partage le composition root de l'adapter Tauri au lieu d'en bâtir un second dans le même processus. `ServerState` porte désormais `Arc<BackendCore>` ; `run_embedded` est conservé en compat et délègue en construisant son propre core. Corrige au passage un vrai bug produit, découvert en rejouant les tests hors sandbox : `run_embedded` ne réconciliait pas le port effectif avec la config. Sur un bind éphémère (`127.0.0.1:0`), l'état du serveur gardait `listen` à port 0, donc `origin_allowed` comparait l'Origin entrante à `http://127.0.0.1:0` et le serveur embarqué rejetait **toutes** les requêtes API en 403. La config reçoit maintenant `local_addr` lu sur le listener **avant** la construction du state. Comble le trou signalé au merge de #65 : `run_embedded().stop()` est enfin couvert. ATTENTION — le premier « 57 passed » de ce lot était un FAUX VERT. Le sandbox d'exécution interdit `TcpListener::bind` ; les tests sortaient en silence sur EPERM, dont précisément le test `stop()`. Les replis EPERM sont supprimés : un bind refusé fait désormais échouer le test au lieu de le peindre en vert. QA ré-exécutée par Git HORS SANDBOX avant merge (sinon on reproduit le faux vert) : `cargo test -p web-server` 57 passed, 0 échec, avec `run_embedded_stop_shuts_down_accept_loop` et `run_embedded_with_core_uses_injected_core_for_http_invokes` réellement exécutés sur le vrai chemin réseau. `-p backend` 41 · `-p app-tauri` 35, 0 échec. DETTE CONNUE, tracée en #72 B0 : le chemin CLI standalone `run_server` garde la même hypothèse fausse — il reçoit un `ServerState` construit AVANT le bind, donc `idea-serve --listen 127.0.0.1:0` reste cassé à l'identique. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
4312 lines
137 KiB
Rust
4312 lines
137 KiB
Rust
//! Secure HTTP driving adapter for `idea --serve`.
|
|
//!
|
|
//! The shared backend core stays unaware of HTTP, cookies, origins and
|
|
//! WebSocket framing; this module owns the secure web driving adapter.
|
|
|
|
use std::collections::HashSet;
|
|
use std::env;
|
|
use std::net::SocketAddr;
|
|
use std::path::{Path, PathBuf};
|
|
use std::process::ExitCode;
|
|
use std::sync::atomic::{AtomicU64, Ordering};
|
|
use std::sync::{Arc, Mutex};
|
|
|
|
use backend::stream::{OutputBridge, OutputSink, OutputSinkError};
|
|
use bytes::Bytes;
|
|
use cookie::{Cookie, SameSite};
|
|
use domain::events::DomainEvent;
|
|
use domain::ports::BackgroundTaskPortError;
|
|
use http::header::{HeaderValue, CONTENT_TYPE, COOKIE, ORIGIN, SET_COOKIE};
|
|
#[cfg(test)]
|
|
use http::Request;
|
|
use http::{HeaderMap, Method, Response, StatusCode, Uri};
|
|
use http_body_util::{BodyExt, Full};
|
|
use serde::{Deserialize, Serialize};
|
|
use serde_json::{json, Value};
|
|
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt};
|
|
use tokio::net::TcpListener;
|
|
use tokio::sync::{mpsc, oneshot};
|
|
use tokio::task::JoinHandle;
|
|
use uuid::Uuid;
|
|
|
|
use application::{
|
|
CloseTerminalInput, GetProjectWorkStateInput, LaunchAgentInput, McpRuntime, OpenProjectInput,
|
|
ResizeTerminalInput, RotateConversationLogInput, WriteToTerminalInput,
|
|
};
|
|
use domain::ports::PtyHandle;
|
|
use domain::{Project, SessionId};
|
|
|
|
use backend::dto::{
|
|
parse_agent_id, parse_node_id, parse_project_id, parse_session_id, parse_task_id,
|
|
BackgroundTaskDto, ErrorDto, HealthRequestDto, HealthResponseDto, LaunchAgentRequestDto,
|
|
OpenTerminalRequestDto, ProjectDto, ProjectListDto, ProjectWorkStateDto, TerminalSessionDto,
|
|
};
|
|
use backend::events::DomainEventDto;
|
|
type PtyChunk = Vec<u8>;
|
|
use backend::BackendCore;
|
|
|
|
const DEFAULT_LISTEN: &str = "127.0.0.1:17373";
|
|
const SESSION_COOKIE: &str = "idea_session";
|
|
const WS_PATH: &str = "/api/ws";
|
|
const WS_MAGIC: &str = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
|
|
const WS_MAX_PAYLOAD: usize = 64 * 1024;
|
|
const WS_OUTPUT_BUFFER: usize = 512;
|
|
|
|
type ResponseBody = Full<Bytes>;
|
|
|
|
/// Runs the `idea --serve` subcommand from already-split CLI arguments.
|
|
pub fn run_from_args(args: Vec<String>) -> ExitCode {
|
|
let config = match ServerConfig::from_args(args).and_then(|config| {
|
|
config.validate()?;
|
|
Ok(config)
|
|
}) {
|
|
Ok(config) => config,
|
|
Err(err) => {
|
|
eprintln!("idea --serve: {err}");
|
|
return ExitCode::from(2);
|
|
}
|
|
};
|
|
|
|
let state = Arc::new(ServerState::new(config.clone()));
|
|
eprintln!("IdeA pairing code: {}", state.pairing_code());
|
|
eprintln!(
|
|
"idea --serve: app data dir = {}",
|
|
config.app_data_dir.display()
|
|
);
|
|
eprintln!("IdeA server listening on {}", config.listen);
|
|
|
|
match tokio::runtime::Builder::new_multi_thread()
|
|
.enable_all()
|
|
.build()
|
|
{
|
|
Ok(runtime) => match runtime.block_on(run_server(config, state)) {
|
|
Ok(()) => ExitCode::SUCCESS,
|
|
Err(err) => {
|
|
eprintln!("idea --serve: {err}");
|
|
ExitCode::from(1)
|
|
}
|
|
},
|
|
Err(err) => {
|
|
eprintln!("idea --serve: failed to build runtime: {err}");
|
|
ExitCode::from(1)
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Configuration shared by the standalone and embedded web-server entry points.
|
|
#[derive(Clone, Debug)]
|
|
pub struct ServerConfig {
|
|
/// Address to bind.
|
|
pub listen: SocketAddr,
|
|
/// Exact public origin allowed by CORS/origin checks.
|
|
pub public_origin: Option<String>,
|
|
/// Whether non-loopback binds are allowed.
|
|
pub allow_remote: bool,
|
|
/// Whether a trusted HTTPS reverse proxy terminates remote access.
|
|
pub trust_reverse_proxy: bool,
|
|
/// IdeA application data directory.
|
|
pub app_data_dir: PathBuf,
|
|
/// Built frontend assets root.
|
|
pub web_root: PathBuf,
|
|
}
|
|
|
|
impl ServerConfig {
|
|
/// Parses the web-server CLI flags.
|
|
pub fn from_args(args: Vec<String>) -> Result<Self, String> {
|
|
let mut listen = DEFAULT_LISTEN
|
|
.parse::<SocketAddr>()
|
|
.expect("default listen address is valid");
|
|
let mut public_origin = None;
|
|
let mut allow_remote = false;
|
|
let mut trust_reverse_proxy = false;
|
|
let mut app_data_dir = default_app_data_dir();
|
|
let mut web_root = None;
|
|
|
|
let mut it = args.into_iter();
|
|
while let Some(arg) = it.next() {
|
|
match arg.as_str() {
|
|
"--listen" => {
|
|
let value = it
|
|
.next()
|
|
.ok_or_else(|| "--listen requires an address".to_owned())?;
|
|
listen = value
|
|
.parse()
|
|
.map_err(|_| format!("invalid --listen address: {value}"))?;
|
|
}
|
|
"--public-origin" => {
|
|
let value = it
|
|
.next()
|
|
.ok_or_else(|| "--public-origin requires an origin".to_owned())?;
|
|
public_origin = Some(value);
|
|
}
|
|
"--allow-remote" => allow_remote = true,
|
|
"--trust-reverse-proxy" => trust_reverse_proxy = true,
|
|
"--app-data-dir" => {
|
|
let value = it
|
|
.next()
|
|
.ok_or_else(|| "--app-data-dir requires a path".to_owned())?;
|
|
app_data_dir = PathBuf::from(value);
|
|
}
|
|
"--web-root" => {
|
|
let value = it
|
|
.next()
|
|
.ok_or_else(|| "--web-root requires a path".to_owned())?;
|
|
web_root = Some(PathBuf::from(value));
|
|
}
|
|
"--help" | "-h" => return Err(Self::usage()),
|
|
other => return Err(format!("unknown --serve argument: {other}")),
|
|
}
|
|
}
|
|
let web_root = resolve_web_root(web_root)?;
|
|
|
|
Ok(Self {
|
|
listen,
|
|
public_origin,
|
|
allow_remote,
|
|
trust_reverse_proxy,
|
|
app_data_dir,
|
|
web_root,
|
|
})
|
|
}
|
|
|
|
/// Validates security-sensitive configuration.
|
|
pub fn validate(&self) -> Result<(), String> {
|
|
if let Some(origin) = &self.public_origin {
|
|
validate_origin(origin)?;
|
|
}
|
|
|
|
if !self.listen.ip().is_loopback() && !self.allow_remote {
|
|
return Err(
|
|
"refusing non-loopback bind without --allow-remote and HTTPS proxy config"
|
|
.to_owned(),
|
|
);
|
|
}
|
|
|
|
if self.allow_remote {
|
|
let Some(origin) = &self.public_origin else {
|
|
return Err("--allow-remote requires --public-origin https://...".to_owned());
|
|
};
|
|
if !origin.starts_with("https://") {
|
|
return Err("--allow-remote requires an HTTPS public origin".to_owned());
|
|
}
|
|
if !self.trust_reverse_proxy {
|
|
return Err("--allow-remote requires --trust-reverse-proxy".to_owned());
|
|
}
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
fn secure_cookie(&self) -> bool {
|
|
self.allow_remote
|
|
|| self
|
|
.public_origin
|
|
.as_deref()
|
|
.is_some_and(|o| o.starts_with("https://"))
|
|
}
|
|
|
|
/// Human-readable CLI usage.
|
|
#[must_use]
|
|
pub fn usage() -> String {
|
|
"usage: idea-serve [--listen IP:PORT] [--app-data-dir PATH] [--web-root PATH] [--allow-remote --public-origin https://host --trust-reverse-proxy]".to_owned()
|
|
}
|
|
}
|
|
|
|
/// Current lifecycle state of an embedded web server.
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
pub enum EmbeddedServerState {
|
|
/// The accept loop has been spawned.
|
|
Running,
|
|
/// A shutdown signal was sent.
|
|
Stopping,
|
|
}
|
|
|
|
/// Handle returned by [`run_embedded`] and [`run_embedded_with_core`].
|
|
pub struct EmbeddedServerHandle {
|
|
url: String,
|
|
pairing_code: String,
|
|
shutdown: Option<oneshot::Sender<()>>,
|
|
task: JoinHandle<Result<(), String>>,
|
|
state: EmbeddedServerState,
|
|
}
|
|
|
|
impl EmbeddedServerHandle {
|
|
/// Effective base URL after binding. Includes the OS-assigned port when
|
|
/// `listen` used port `0`.
|
|
#[must_use]
|
|
pub fn url(&self) -> &str {
|
|
&self.url
|
|
}
|
|
|
|
/// Pairing code accepted by `POST /api/pair`.
|
|
#[must_use]
|
|
pub fn pairing_code(&self) -> &str {
|
|
&self.pairing_code
|
|
}
|
|
|
|
/// Current lifecycle state.
|
|
#[must_use]
|
|
pub fn state(&self) -> EmbeddedServerState {
|
|
self.state
|
|
}
|
|
|
|
/// Requests shutdown and waits for the accept loop to finish.
|
|
pub async fn stop(mut self) -> Result<(), String> {
|
|
self.state = EmbeddedServerState::Stopping;
|
|
if let Some(shutdown) = self.shutdown.take() {
|
|
let _ = shutdown.send(());
|
|
}
|
|
self.task
|
|
.await
|
|
.map_err(|err| format!("embedded web server task failed: {err}"))?
|
|
}
|
|
}
|
|
|
|
/// Starts the web server as an embedded library component.
|
|
///
|
|
/// This compatibility entry point builds its own backend composition root. The
|
|
/// desktop app should use [`run_embedded_with_core`] so HTTP handlers share the
|
|
/// already-running backend core.
|
|
pub async fn run_embedded(config: ServerConfig) -> Result<EmbeddedServerHandle, String> {
|
|
let core = Arc::new(BackendCore::build(config.app_data_dir.clone()));
|
|
run_embedded_with_core(config, core).await
|
|
}
|
|
|
|
/// Starts the web server with an already-built backend core.
|
|
///
|
|
/// This is the integration point for the desktop-hosted server: the HTTP
|
|
/// driving adapter receives the same composition root as the Tauri adapter
|
|
/// instead of constructing a second backend core in the same process.
|
|
pub async fn run_embedded_with_core(
|
|
config: ServerConfig,
|
|
core: Arc<BackendCore>,
|
|
) -> Result<EmbeddedServerHandle, String> {
|
|
config.validate()?;
|
|
let listener = TcpListener::bind(config.listen)
|
|
.await
|
|
.map_err(|err| format!("failed to bind {}: {err}", config.listen))?;
|
|
let local_addr = listener
|
|
.local_addr()
|
|
.map_err(|err| format!("failed to read listener address: {err}"))?;
|
|
let mut effective_config = config.clone();
|
|
effective_config.listen = local_addr;
|
|
let state = Arc::new(ServerState::with_core(effective_config, core));
|
|
let pairing_code = state.pairing_code().to_owned();
|
|
let (shutdown_tx, shutdown_rx) = oneshot::channel();
|
|
let task = tokio::spawn(run_listener(listener, state, shutdown_rx));
|
|
Ok(EmbeddedServerHandle {
|
|
url: format!("http://{local_addr}"),
|
|
pairing_code,
|
|
shutdown: Some(shutdown_tx),
|
|
task,
|
|
state: EmbeddedServerState::Running,
|
|
})
|
|
}
|
|
|
|
fn resolve_web_root(explicit: Option<PathBuf>) -> Result<PathBuf, String> {
|
|
if let Some(path) = explicit {
|
|
return validate_web_root(path, "--web-root");
|
|
}
|
|
if let Some(path) = env::var_os("IDEA_WEB_ROOT").map(PathBuf::from) {
|
|
return validate_web_root(path, "IDEA_WEB_ROOT");
|
|
}
|
|
|
|
let mut candidates = Vec::new();
|
|
if let Ok(exe) = env::current_exe() {
|
|
if let Some(dir) = exe.parent() {
|
|
candidates.push(dir.join("web"));
|
|
candidates.push(dir.join("frontend").join("dist"));
|
|
}
|
|
}
|
|
if let Ok(cwd) = env::current_dir() {
|
|
candidates.push(cwd.join("frontend").join("dist"));
|
|
}
|
|
|
|
for candidate in candidates {
|
|
if candidate.join("index.html").is_file() {
|
|
return Ok(candidate);
|
|
}
|
|
}
|
|
|
|
Err("web assets not found: build frontend/dist or pass --web-root PATH".to_owned())
|
|
}
|
|
|
|
fn validate_web_root(path: PathBuf, source: &str) -> Result<PathBuf, String> {
|
|
if path.join("index.html").is_file() {
|
|
Ok(path)
|
|
} else {
|
|
Err(format!(
|
|
"{source} does not contain an index.html: {}",
|
|
path.display()
|
|
))
|
|
}
|
|
}
|
|
|
|
struct ServerState {
|
|
config: ServerConfig,
|
|
app: Arc<BackendCore>,
|
|
pairing_code: String,
|
|
sessions: Mutex<HashSet<String>>,
|
|
ws_pty_bridge: Arc<OutputBridge<SessionId, PtyChunk>>,
|
|
security_logger: Arc<dyn SecurityLogger>,
|
|
}
|
|
|
|
impl ServerState {
|
|
fn new(config: ServerConfig) -> Self {
|
|
let core = Arc::new(BackendCore::build(config.app_data_dir.clone()));
|
|
Self::with_core(config, core)
|
|
}
|
|
|
|
fn with_core(config: ServerConfig, core: Arc<BackendCore>) -> Self {
|
|
Self {
|
|
app: core,
|
|
config,
|
|
pairing_code: new_pairing_code(),
|
|
sessions: Mutex::new(HashSet::new()),
|
|
ws_pty_bridge: Arc::new(OutputBridge::new()),
|
|
security_logger: Arc::new(StderrSecurityLogger),
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
fn new_for_test(config: ServerConfig, pairing_code: impl Into<String>) -> Self {
|
|
Self::new_for_test_with_logger(config, pairing_code, Arc::new(NoopSecurityLogger))
|
|
}
|
|
|
|
#[cfg(test)]
|
|
fn new_for_test_with_logger(
|
|
config: ServerConfig,
|
|
pairing_code: impl Into<String>,
|
|
security_logger: Arc<dyn SecurityLogger>,
|
|
) -> Self {
|
|
Self {
|
|
app: Arc::new(BackendCore::build(config.app_data_dir.clone())),
|
|
config,
|
|
pairing_code: pairing_code.into(),
|
|
sessions: Mutex::new(HashSet::new()),
|
|
ws_pty_bridge: Arc::new(OutputBridge::new()),
|
|
security_logger,
|
|
}
|
|
}
|
|
|
|
fn pairing_code(&self) -> &str {
|
|
&self.pairing_code
|
|
}
|
|
|
|
fn create_session(&self) -> String {
|
|
let token = new_session_token();
|
|
if let Ok(mut sessions) = self.sessions.lock() {
|
|
sessions.insert(token.clone());
|
|
}
|
|
token
|
|
}
|
|
|
|
fn has_session(&self, token: &str) -> bool {
|
|
self.sessions
|
|
.lock()
|
|
.map(|sessions| sessions.contains(token))
|
|
.unwrap_or(false)
|
|
}
|
|
|
|
fn revoke_session(&self, token: &str) -> bool {
|
|
self.sessions
|
|
.lock()
|
|
.map(|mut sessions| sessions.remove(token))
|
|
.unwrap_or(false)
|
|
}
|
|
|
|
fn log_security(&self, event: SecurityLogEvent) {
|
|
self.security_logger.log(event);
|
|
}
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
enum SecurityLogEvent {
|
|
PairingSucceeded {
|
|
origin: Option<String>,
|
|
},
|
|
PairingFailed {
|
|
origin: Option<String>,
|
|
reason: &'static str,
|
|
},
|
|
OriginRejected {
|
|
origin: Option<String>,
|
|
route: String,
|
|
},
|
|
WsUpgradeRejected {
|
|
origin: Option<String>,
|
|
reason: &'static str,
|
|
},
|
|
SessionRevoked {
|
|
origin: Option<String>,
|
|
},
|
|
}
|
|
|
|
trait SecurityLogger: Send + Sync {
|
|
fn log(&self, event: SecurityLogEvent);
|
|
}
|
|
|
|
struct StderrSecurityLogger;
|
|
|
|
impl SecurityLogger for StderrSecurityLogger {
|
|
fn log(&self, event: SecurityLogEvent) {
|
|
eprintln!("idea --serve security: {}", security_log_line(&event));
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
struct NoopSecurityLogger;
|
|
|
|
#[cfg(test)]
|
|
impl SecurityLogger for NoopSecurityLogger {
|
|
fn log(&self, _event: SecurityLogEvent) {}
|
|
}
|
|
|
|
fn security_log_line(event: &SecurityLogEvent) -> String {
|
|
match event {
|
|
SecurityLogEvent::PairingSucceeded { origin } => {
|
|
format!("pairing succeeded origin={}", origin_label(origin))
|
|
}
|
|
SecurityLogEvent::PairingFailed { origin, reason } => {
|
|
format!(
|
|
"pairing failed reason={reason} origin={}",
|
|
origin_label(origin)
|
|
)
|
|
}
|
|
SecurityLogEvent::OriginRejected { origin, route } => {
|
|
format!(
|
|
"origin rejected route={} origin={}",
|
|
route,
|
|
origin_label(origin)
|
|
)
|
|
}
|
|
SecurityLogEvent::WsUpgradeRejected { origin, reason } => {
|
|
format!(
|
|
"websocket upgrade rejected reason={reason} origin={}",
|
|
origin_label(origin)
|
|
)
|
|
}
|
|
SecurityLogEvent::SessionRevoked { origin } => {
|
|
format!("session revoked origin={}", origin_label(origin))
|
|
}
|
|
}
|
|
}
|
|
|
|
fn origin_label(origin: &Option<String>) -> &str {
|
|
origin.as_deref().unwrap_or("<missing>")
|
|
}
|
|
|
|
async fn run_server(config: ServerConfig, state: Arc<ServerState>) -> Result<(), String> {
|
|
let listener = TcpListener::bind(config.listen)
|
|
.await
|
|
.map_err(|err| format!("failed to bind {}: {err}", config.listen))?;
|
|
let (_shutdown_tx, shutdown_rx) = oneshot::channel();
|
|
run_listener(listener, state, shutdown_rx).await
|
|
}
|
|
|
|
async fn run_listener(
|
|
listener: TcpListener,
|
|
state: Arc<ServerState>,
|
|
mut shutdown: oneshot::Receiver<()>,
|
|
) -> Result<(), String> {
|
|
loop {
|
|
let accepted = tokio::select! {
|
|
result = listener.accept() => result,
|
|
_ = &mut shutdown => return Ok(()),
|
|
};
|
|
let (stream, _) = accepted.map_err(|err| format!("failed to accept connection: {err}"))?;
|
|
let state = Arc::clone(&state);
|
|
tokio::spawn(async move {
|
|
if let Err(err) = handle_tcp_connection(stream, state).await {
|
|
eprintln!("idea --serve: connection error: {err}");
|
|
}
|
|
});
|
|
}
|
|
}
|
|
|
|
async fn handle_tcp_connection(
|
|
mut stream: tokio::net::TcpStream,
|
|
state: Arc<ServerState>,
|
|
) -> Result<(), String> {
|
|
let mut buffer = Vec::with_capacity(8192);
|
|
let mut chunk = [0_u8; 2048];
|
|
let header_end = loop {
|
|
let read = stream
|
|
.read(&mut chunk)
|
|
.await
|
|
.map_err(|err| format!("failed to read request: {err}"))?;
|
|
if read == 0 {
|
|
return Ok(());
|
|
}
|
|
buffer.extend_from_slice(&chunk[..read]);
|
|
if buffer.len() > 1024 * 1024 {
|
|
return Err("request too large".to_owned());
|
|
}
|
|
if let Some(pos) = find_header_end(&buffer) {
|
|
break pos;
|
|
}
|
|
};
|
|
|
|
let (method, uri, headers, content_length) = parse_http_request_head(&buffer[..header_end])?;
|
|
if method == Method::GET && uri.path() == WS_PATH {
|
|
return handle_ws_upgrade(stream, headers, state).await;
|
|
}
|
|
|
|
let body_start = header_end + 4;
|
|
while buffer.len() < body_start + content_length {
|
|
let read = stream
|
|
.read(&mut chunk)
|
|
.await
|
|
.map_err(|err| format!("failed to read request body: {err}"))?;
|
|
if read == 0 {
|
|
break;
|
|
}
|
|
buffer.extend_from_slice(&chunk[..read]);
|
|
if buffer.len() > 1024 * 1024 {
|
|
return Err("request too large".to_owned());
|
|
}
|
|
}
|
|
if buffer.len() < body_start + content_length {
|
|
return Err("truncated request body".to_owned());
|
|
}
|
|
|
|
let response = dispatch_http(
|
|
method,
|
|
uri,
|
|
headers,
|
|
Bytes::copy_from_slice(&buffer[body_start..body_start + content_length]),
|
|
state,
|
|
)
|
|
.await;
|
|
write_http_response(&mut stream, response).await
|
|
}
|
|
|
|
#[cfg(test)]
|
|
async fn handle_request(
|
|
req: Request<ResponseBody>,
|
|
state: Arc<ServerState>,
|
|
) -> Response<ResponseBody> {
|
|
let (parts, body) = req.into_parts();
|
|
let body = match body.collect().await {
|
|
Ok(body) => body.to_bytes(),
|
|
Err(err) => {
|
|
return error_response(
|
|
StatusCode::BAD_REQUEST,
|
|
"INVALID",
|
|
format!("invalid request body: {err}"),
|
|
None,
|
|
);
|
|
}
|
|
};
|
|
|
|
dispatch_http(
|
|
parts.method,
|
|
parts.uri,
|
|
parts.headers,
|
|
body,
|
|
Arc::clone(&state),
|
|
)
|
|
.await
|
|
}
|
|
|
|
fn find_header_end(buffer: &[u8]) -> Option<usize> {
|
|
buffer.windows(4).position(|window| window == b"\r\n\r\n")
|
|
}
|
|
|
|
fn parse_http_request_head(head: &[u8]) -> Result<(Method, Uri, HeaderMap, usize), String> {
|
|
let text = std::str::from_utf8(head).map_err(|_| "request head is not UTF-8".to_owned())?;
|
|
let mut lines = text.split("\r\n");
|
|
let request_line = lines
|
|
.next()
|
|
.ok_or_else(|| "missing request line".to_owned())?;
|
|
let mut request_parts = request_line.split_whitespace();
|
|
let method = request_parts
|
|
.next()
|
|
.ok_or_else(|| "missing method".to_owned())?
|
|
.parse::<Method>()
|
|
.map_err(|_| "invalid method".to_owned())?;
|
|
let uri = request_parts
|
|
.next()
|
|
.ok_or_else(|| "missing uri".to_owned())?
|
|
.parse::<Uri>()
|
|
.map_err(|_| "invalid uri".to_owned())?;
|
|
let mut headers = HeaderMap::new();
|
|
let mut content_length = 0;
|
|
|
|
for line in lines {
|
|
if line.is_empty() {
|
|
continue;
|
|
}
|
|
let Some((name, value)) = line.split_once(':') else {
|
|
return Err("invalid header line".to_owned());
|
|
};
|
|
let name = http::header::HeaderName::from_bytes(name.trim().as_bytes())
|
|
.map_err(|_| "invalid header name".to_owned())?;
|
|
let value =
|
|
HeaderValue::from_str(value.trim()).map_err(|_| "invalid header value".to_owned())?;
|
|
if name == http::header::CONTENT_LENGTH {
|
|
content_length = value
|
|
.to_str()
|
|
.ok()
|
|
.and_then(|raw| raw.parse::<usize>().ok())
|
|
.ok_or_else(|| "invalid content-length".to_owned())?;
|
|
}
|
|
headers.insert(name, value);
|
|
}
|
|
|
|
Ok((method, uri, headers, content_length))
|
|
}
|
|
|
|
async fn write_http_response(
|
|
stream: &mut tokio::net::TcpStream,
|
|
response: Response<ResponseBody>,
|
|
) -> Result<(), String> {
|
|
let status = response.status();
|
|
let headers = response.headers().clone();
|
|
let body = response
|
|
.into_body()
|
|
.collect()
|
|
.await
|
|
.map_err(|err| format!("failed to collect response: {err}"))?
|
|
.to_bytes();
|
|
let reason = status.canonical_reason().unwrap_or("Unknown");
|
|
let mut bytes = format!("HTTP/1.1 {} {}\r\n", status.as_u16(), reason).into_bytes();
|
|
for (name, value) in &headers {
|
|
bytes.extend_from_slice(name.as_str().as_bytes());
|
|
bytes.extend_from_slice(b": ");
|
|
bytes.extend_from_slice(value.as_bytes());
|
|
bytes.extend_from_slice(b"\r\n");
|
|
}
|
|
bytes.extend_from_slice(format!("content-length: {}\r\n", body.len()).as_bytes());
|
|
bytes.extend_from_slice(b"connection: close\r\n\r\n");
|
|
bytes.extend_from_slice(&body);
|
|
stream
|
|
.write_all(&bytes)
|
|
.await
|
|
.map_err(|err| format!("failed to write response: {err}"))
|
|
}
|
|
|
|
async fn handle_ws_upgrade(
|
|
mut stream: tokio::net::TcpStream,
|
|
headers: HeaderMap,
|
|
state: Arc<ServerState>,
|
|
) -> Result<(), String> {
|
|
let accept = match validate_ws_upgrade(&headers, &state) {
|
|
Ok(accept) => accept,
|
|
Err(response) => return write_http_response(&mut stream, *response).await,
|
|
};
|
|
|
|
let response = format!(
|
|
"HTTP/1.1 101 Switching Protocols\r\n\
|
|
upgrade: websocket\r\n\
|
|
connection: Upgrade\r\n\
|
|
sec-websocket-accept: {accept}\r\n\r\n"
|
|
);
|
|
stream
|
|
.write_all(response.as_bytes())
|
|
.await
|
|
.map_err(|err| format!("failed to write websocket upgrade: {err}"))?;
|
|
run_ws_connection(stream, state).await
|
|
}
|
|
|
|
fn validate_ws_upgrade(
|
|
headers: &HeaderMap,
|
|
state: &ServerState,
|
|
) -> Result<String, Box<Response<ResponseBody>>> {
|
|
let origin = match validate_request_origin(headers, &state.config) {
|
|
Ok(origin) => origin,
|
|
Err(response) => {
|
|
state.log_security(SecurityLogEvent::OriginRejected {
|
|
origin: request_origin(headers),
|
|
route: WS_PATH.to_owned(),
|
|
});
|
|
return Err(response);
|
|
}
|
|
};
|
|
let Some(token) = session_cookie(headers) else {
|
|
state.log_security(SecurityLogEvent::WsUpgradeRejected {
|
|
origin: origin.clone(),
|
|
reason: "missing_session",
|
|
});
|
|
return Err(Box::new(error_response(
|
|
StatusCode::UNAUTHORIZED,
|
|
"UNAUTHORIZED",
|
|
"missing session cookie",
|
|
origin.as_deref(),
|
|
)));
|
|
};
|
|
if !state.has_session(&token) {
|
|
state.log_security(SecurityLogEvent::WsUpgradeRejected {
|
|
origin: origin.clone(),
|
|
reason: "invalid_session",
|
|
});
|
|
return Err(Box::new(error_response(
|
|
StatusCode::UNAUTHORIZED,
|
|
"UNAUTHORIZED",
|
|
"invalid session cookie",
|
|
origin.as_deref(),
|
|
)));
|
|
}
|
|
if !header_contains_token(headers, "upgrade", "websocket")
|
|
|| !header_contains_token(headers, "connection", "upgrade")
|
|
{
|
|
state.log_security(SecurityLogEvent::WsUpgradeRejected {
|
|
origin: origin.clone(),
|
|
reason: "invalid_upgrade",
|
|
});
|
|
return Err(Box::new(error_response(
|
|
StatusCode::BAD_REQUEST,
|
|
"INVALID",
|
|
"missing websocket upgrade headers",
|
|
origin.as_deref(),
|
|
)));
|
|
}
|
|
let Some(key) = headers
|
|
.get("sec-websocket-key")
|
|
.and_then(|value| value.to_str().ok())
|
|
else {
|
|
state.log_security(SecurityLogEvent::WsUpgradeRejected {
|
|
origin: origin.clone(),
|
|
reason: "missing_key",
|
|
});
|
|
return Err(Box::new(error_response(
|
|
StatusCode::BAD_REQUEST,
|
|
"INVALID",
|
|
"missing Sec-WebSocket-Key",
|
|
origin.as_deref(),
|
|
)));
|
|
};
|
|
Ok(websocket_accept(key))
|
|
}
|
|
|
|
fn header_contains_token(headers: &HeaderMap, name: &str, needle: &str) -> bool {
|
|
headers
|
|
.get(name)
|
|
.and_then(|value| value.to_str().ok())
|
|
.is_some_and(|value| {
|
|
value
|
|
.split(',')
|
|
.any(|token| token.trim().eq_ignore_ascii_case(needle))
|
|
})
|
|
}
|
|
|
|
async fn run_ws_connection(
|
|
stream: tokio::net::TcpStream,
|
|
state: Arc<ServerState>,
|
|
) -> Result<(), String> {
|
|
let (mut reader, mut writer) = stream.into_split();
|
|
let (tx, mut rx) = mpsc::channel::<ServerFrame>(WS_OUTPUT_BUFFER);
|
|
let owned = Arc::new(Mutex::new(Vec::<(SessionId, u64)>::new()));
|
|
let event_relay_task =
|
|
spawn_ws_domain_event_relay(state.app.event_bus.raw_receiver(), tx.clone());
|
|
let writer_task = tokio::spawn(async move {
|
|
while let Some(frame) = rx.recv().await {
|
|
let text = serde_json::to_vec(&frame)
|
|
.map_err(|err| format!("failed to encode server frame: {err}"))?;
|
|
let bytes = encode_ws_frame(WsOpcode::Text, &text);
|
|
writer
|
|
.write_all(&bytes)
|
|
.await
|
|
.map_err(|err| format!("failed to write websocket frame: {err}"))?;
|
|
}
|
|
Ok::<(), String>(())
|
|
});
|
|
|
|
loop {
|
|
let frame = match read_ws_frame(&mut reader).await {
|
|
Ok(frame) => frame,
|
|
Err(err) => {
|
|
let _ = tx
|
|
.send(ServerFrame::error(
|
|
None,
|
|
None,
|
|
"WS_PROTOCOL",
|
|
err.to_string(),
|
|
))
|
|
.await;
|
|
break;
|
|
}
|
|
};
|
|
match frame.opcode {
|
|
WsOpcode::Text => {
|
|
let parsed = serde_json::from_slice::<ClientFrame>(&frame.payload)
|
|
.map_err(|err| format!("invalid client frame JSON: {err}"));
|
|
match parsed {
|
|
Ok(frame) => {
|
|
handle_client_frame(frame, &state, &tx, &owned).await;
|
|
}
|
|
Err(err) => {
|
|
let _ = tx
|
|
.send(ServerFrame::error(None, None, "INVALID", err))
|
|
.await;
|
|
}
|
|
}
|
|
}
|
|
WsOpcode::Ping => {
|
|
let _ = tx.send(ServerFrame::pong()).await;
|
|
}
|
|
WsOpcode::Close => break,
|
|
WsOpcode::Pong => {}
|
|
WsOpcode::Binary => {
|
|
let _ = tx
|
|
.send(ServerFrame::error(
|
|
None,
|
|
None,
|
|
"UNSUPPORTED",
|
|
"binary websocket frames are not supported",
|
|
))
|
|
.await;
|
|
}
|
|
}
|
|
}
|
|
|
|
if let Ok(owned) = owned.lock() {
|
|
for (session, gen) in owned.iter() {
|
|
state.ws_pty_bridge.unregister_if(session, *gen);
|
|
}
|
|
}
|
|
event_relay_task.abort();
|
|
drop(tx);
|
|
writer_task
|
|
.await
|
|
.map_err(|err| format!("websocket writer task failed: {err}"))?
|
|
}
|
|
|
|
fn spawn_ws_domain_event_relay(
|
|
mut rx: tokio::sync::broadcast::Receiver<DomainEvent>,
|
|
tx: mpsc::Sender<ServerFrame>,
|
|
) -> tokio::task::JoinHandle<()> {
|
|
use tokio::sync::broadcast::error::RecvError;
|
|
|
|
tokio::spawn(async move {
|
|
loop {
|
|
match rx.recv().await {
|
|
Ok(event) => {
|
|
if matches!(event, DomainEvent::PtyOutput { .. }) {
|
|
continue;
|
|
}
|
|
match tx.try_send(ServerFrame::domain_event(&event)) {
|
|
Ok(()) | Err(mpsc::error::TrySendError::Full(_)) => {}
|
|
Err(mpsc::error::TrySendError::Closed(_)) => break,
|
|
}
|
|
}
|
|
Err(RecvError::Lagged(_)) => continue,
|
|
Err(RecvError::Closed) => break,
|
|
}
|
|
}
|
|
})
|
|
}
|
|
|
|
async fn handle_client_frame(
|
|
frame: ClientFrame,
|
|
state: &Arc<ServerState>,
|
|
tx: &mpsc::Sender<ServerFrame>,
|
|
owned: &Arc<Mutex<Vec<(SessionId, u64)>>>,
|
|
) {
|
|
let result = match frame.kind.as_str() {
|
|
"terminal.open" | "open_terminal" => ws_open_terminal(&frame, state, tx, owned).await,
|
|
"agent.launch" | "launch_agent" => ws_launch_agent(&frame, state, tx, owned).await,
|
|
"terminal.attach" | "attach_terminal" => ws_attach_terminal(&frame, state, tx, owned).await,
|
|
"terminal.input" | "input" => ws_input(&frame, state),
|
|
"terminal.resize" | "resize" => ws_resize(&frame, state),
|
|
"terminal.detach" | "detach" => ws_detach(&frame, state, owned),
|
|
"terminal.close" | "close" => match ws_close(&frame, state, owned).await {
|
|
Ok((session_id, exit_code)) => {
|
|
let _ = tx
|
|
.send(ServerFrame::status(session_id, "exited", exit_code))
|
|
.await;
|
|
Ok(())
|
|
}
|
|
Err(err) => Err(err),
|
|
},
|
|
"ping" => {
|
|
let _ = tx.send(ServerFrame::pong()).await;
|
|
Ok(())
|
|
}
|
|
_ => Err(ErrorDto {
|
|
code: "UNKNOWN_FRAME".to_owned(),
|
|
message: format!("unknown websocket frame kind: {}", frame.kind),
|
|
}),
|
|
};
|
|
|
|
if let Err(err) = result {
|
|
let _ = tx
|
|
.send(ServerFrame::error(
|
|
Some(frame.id),
|
|
payload_session_id(&frame.payload),
|
|
err.code,
|
|
err.message,
|
|
))
|
|
.await;
|
|
}
|
|
}
|
|
|
|
async fn ws_open_terminal(
|
|
frame: &ClientFrame,
|
|
state: &Arc<ServerState>,
|
|
tx: &mpsc::Sender<ServerFrame>,
|
|
owned: &Arc<Mutex<Vec<(SessionId, u64)>>>,
|
|
) -> Result<(), ErrorDto> {
|
|
let request_value = frame
|
|
.payload
|
|
.get("request")
|
|
.cloned()
|
|
.unwrap_or_else(|| frame.payload.clone());
|
|
let request: OpenTerminalRequestDto =
|
|
serde_json::from_value(request_value).map_err(invalid_args_error)?;
|
|
let output = state
|
|
.app
|
|
.open_terminal
|
|
.execute(request.into())
|
|
.await
|
|
.map_err(ErrorDto::from)?;
|
|
let dto = TerminalSessionDto::from(output);
|
|
let sid = parse_session_id(&dto.session_id)?;
|
|
let sink = WsPtySink::new(sid, tx.clone(), 0);
|
|
tx.send(ServerFrame::attached(&frame.id, &dto, Vec::new(), 0, false))
|
|
.await
|
|
.map_err(|_| ErrorDto {
|
|
code: "WS_CLOSED".to_owned(),
|
|
message: "websocket output closed".to_owned(),
|
|
})?;
|
|
attach_sink_and_pump(state, sid, sink, owned)
|
|
}
|
|
|
|
async fn ws_launch_agent(
|
|
frame: &ClientFrame,
|
|
state: &Arc<ServerState>,
|
|
tx: &mpsc::Sender<ServerFrame>,
|
|
owned: &Arc<Mutex<Vec<(SessionId, u64)>>>,
|
|
) -> Result<(), ErrorDto> {
|
|
let request_value = frame
|
|
.payload
|
|
.get("request")
|
|
.cloned()
|
|
.unwrap_or_else(|| frame.payload.clone());
|
|
let request: LaunchAgentRequestDto =
|
|
serde_json::from_value(request_value).map_err(invalid_args_error)?;
|
|
let output = execute_launch_agent_for_ws(state, request).await?;
|
|
send_launch_agent_attached(frame, state, tx, owned, output).await
|
|
}
|
|
|
|
async fn send_launch_agent_attached(
|
|
frame: &ClientFrame,
|
|
state: &Arc<ServerState>,
|
|
tx: &mpsc::Sender<ServerFrame>,
|
|
owned: &Arc<Mutex<Vec<(SessionId, u64)>>>,
|
|
output: application::LaunchAgentOutput,
|
|
) -> Result<(), ErrorDto> {
|
|
if output.structured.is_some() {
|
|
return Err(ErrorDto {
|
|
code: "UNSUPPORTED".to_owned(),
|
|
message: "structured agent sessions do not stream over the PTY websocket".to_owned(),
|
|
});
|
|
}
|
|
let session_id = output.session.id;
|
|
let dto = TerminalSessionDto::from(output);
|
|
let sid = parse_session_id(&dto.session_id)?;
|
|
let sink = WsPtySink::new(sid, tx.clone(), 0);
|
|
tx.send(ServerFrame::attached(&frame.id, &dto, Vec::new(), 0, false))
|
|
.await
|
|
.map_err(|_| ErrorDto {
|
|
code: "WS_CLOSED".to_owned(),
|
|
message: "websocket output closed".to_owned(),
|
|
})?;
|
|
attach_sink_and_pump(state, session_id, sink, owned)
|
|
}
|
|
|
|
async fn execute_launch_agent_for_ws(
|
|
state: &Arc<ServerState>,
|
|
request: LaunchAgentRequestDto,
|
|
) -> Result<application::LaunchAgentOutput, ErrorDto> {
|
|
let project = resolve_project_readonly(&request.project_id, &state.app).await?;
|
|
let agent_id = parse_agent_id(&request.agent_id)?;
|
|
let node_id = request.node_id.as_deref().map(parse_node_id).transpose()?;
|
|
let mcp_runtime = backend::mcp_endpoint::idea_exe_path().map(|exe| McpRuntime {
|
|
exe,
|
|
endpoint: backend::mcp_endpoint::mcp_endpoint(&project.id)
|
|
.as_cli_arg()
|
|
.to_owned(),
|
|
project_id: project.id.as_uuid().simple().to_string(),
|
|
requester: agent_id.to_string(),
|
|
});
|
|
|
|
state.app.reconcile_claude_run_dirs(&project).await;
|
|
|
|
let resume_project = project.clone();
|
|
let rotation_root = project.root.clone();
|
|
let watch_root = project.root.clone();
|
|
let output = state
|
|
.app
|
|
.launch_agent
|
|
.execute(LaunchAgentInput {
|
|
project,
|
|
agent_id,
|
|
rows: request.rows,
|
|
cols: request.cols,
|
|
node_id,
|
|
conversation_id: request.conversation_id.clone(),
|
|
mcp_runtime,
|
|
allow_structured_alongside_pty: false,
|
|
})
|
|
.await
|
|
.map_err(ErrorDto::from)?;
|
|
|
|
if let Ok(mut contexts) = state.app.resume_contexts.lock() {
|
|
contexts.insert(
|
|
agent_id,
|
|
backend::ResumeContext {
|
|
project: resume_project,
|
|
rows: request.rows,
|
|
cols: request.cols,
|
|
},
|
|
);
|
|
}
|
|
|
|
if let Some(profile) = output.profile.as_ref() {
|
|
state.app.arm_turn_watch(
|
|
&watch_root,
|
|
agent_id,
|
|
profile,
|
|
output.assigned_conversation_id.clone(),
|
|
);
|
|
}
|
|
|
|
if let Some(conversation) = request
|
|
.conversation_id
|
|
.as_deref()
|
|
.and_then(|raw| uuid::Uuid::parse_str(raw).ok())
|
|
.map(domain::ConversationId::from_uuid)
|
|
{
|
|
let rotate = Arc::clone(&state.app.rotate_conversation_log);
|
|
tokio::spawn(async move {
|
|
let _ = rotate
|
|
.execute(RotateConversationLogInput {
|
|
project_root: rotation_root,
|
|
conversation,
|
|
})
|
|
.await;
|
|
});
|
|
}
|
|
|
|
Ok(output)
|
|
}
|
|
|
|
async fn ws_attach_terminal(
|
|
frame: &ClientFrame,
|
|
state: &Arc<ServerState>,
|
|
tx: &mpsc::Sender<ServerFrame>,
|
|
owned: &Arc<Mutex<Vec<(SessionId, u64)>>>,
|
|
) -> Result<(), ErrorDto> {
|
|
let session_id = payload_string(&frame.payload, "sessionId")?;
|
|
let sid = parse_session_id(&session_id)?;
|
|
let handle = PtyHandle { session_id: sid };
|
|
let scrollback = state
|
|
.app
|
|
.pty_port
|
|
.scrollback(&handle)
|
|
.map_err(|err| ErrorDto::from(application::AppError::from(err)))?;
|
|
let rows = payload_u16(&frame.payload, "rows").unwrap_or(24);
|
|
let cols = payload_u16(&frame.payload, "cols").unwrap_or(80);
|
|
let dto = TerminalSessionDto {
|
|
session_id,
|
|
cwd: String::new(),
|
|
rows,
|
|
cols,
|
|
assigned_conversation_id: None,
|
|
engine_session_id: None,
|
|
cell_kind: backend::dto::CellKind::Pty,
|
|
};
|
|
let next_seq = u64::from(!scrollback.is_empty());
|
|
tx.send(ServerFrame::attached(
|
|
&frame.id,
|
|
&dto,
|
|
scrollback.clone(),
|
|
next_seq,
|
|
frame.payload.get("lastSeq").is_some(),
|
|
))
|
|
.await
|
|
.map_err(|_| ErrorDto {
|
|
code: "WS_CLOSED".to_owned(),
|
|
message: "websocket output closed".to_owned(),
|
|
})?;
|
|
let sink = WsPtySink::new(sid, tx.clone(), next_seq);
|
|
attach_sink_and_pump(state, sid, sink, owned)
|
|
}
|
|
|
|
fn attach_sink_and_pump(
|
|
state: &Arc<ServerState>,
|
|
sid: SessionId,
|
|
sink: WsPtySink,
|
|
owned: &Arc<Mutex<Vec<(SessionId, u64)>>>,
|
|
) -> Result<(), ErrorDto> {
|
|
let gen = state.ws_pty_bridge.register(sid, Arc::new(sink));
|
|
if let Ok(mut owned) = owned.lock() {
|
|
owned.push((sid, gen));
|
|
}
|
|
let handle = PtyHandle { session_id: sid };
|
|
let stream = state
|
|
.app
|
|
.pty_port
|
|
.subscribe_output(&handle)
|
|
.map_err(|err| ErrorDto::from(application::AppError::from(err)))?;
|
|
let bridge = Arc::clone(&state.ws_pty_bridge);
|
|
std::thread::spawn(move || {
|
|
for chunk in stream {
|
|
if !bridge.send_output(&sid, chunk) {
|
|
break;
|
|
}
|
|
}
|
|
bridge.unregister_if(&sid, gen);
|
|
});
|
|
Ok(())
|
|
}
|
|
|
|
fn ws_input(frame: &ClientFrame, state: &Arc<ServerState>) -> Result<(), ErrorDto> {
|
|
let sid = parse_session_id(&payload_string(&frame.payload, "sessionId")?)?;
|
|
let bytes = payload_bytes(&frame.payload)?;
|
|
state
|
|
.app
|
|
.write_terminal
|
|
.execute(WriteToTerminalInput {
|
|
session_id: sid,
|
|
data: bytes,
|
|
})
|
|
.map_err(ErrorDto::from)
|
|
}
|
|
|
|
fn ws_resize(frame: &ClientFrame, state: &Arc<ServerState>) -> Result<(), ErrorDto> {
|
|
let sid = parse_session_id(&payload_string(&frame.payload, "sessionId")?)?;
|
|
let rows = payload_u16(&frame.payload, "rows").ok_or_else(|| ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: "resize requires rows".to_owned(),
|
|
})?;
|
|
let cols = payload_u16(&frame.payload, "cols").ok_or_else(|| ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: "resize requires cols".to_owned(),
|
|
})?;
|
|
state
|
|
.app
|
|
.resize_terminal
|
|
.execute(ResizeTerminalInput {
|
|
session_id: sid,
|
|
rows,
|
|
cols,
|
|
})
|
|
.map_err(ErrorDto::from)
|
|
}
|
|
|
|
fn ws_detach(
|
|
frame: &ClientFrame,
|
|
state: &Arc<ServerState>,
|
|
owned: &Arc<Mutex<Vec<(SessionId, u64)>>>,
|
|
) -> Result<(), ErrorDto> {
|
|
let sid = parse_session_id(&payload_string(&frame.payload, "sessionId")?)?;
|
|
if let Ok(mut owned) = owned.lock() {
|
|
if let Some(index) = owned.iter().rposition(|(session, _)| *session == sid) {
|
|
let (_, gen) = owned.remove(index);
|
|
state.ws_pty_bridge.unregister_if(&sid, gen);
|
|
}
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
async fn ws_close(
|
|
frame: &ClientFrame,
|
|
state: &Arc<ServerState>,
|
|
owned: &Arc<Mutex<Vec<(SessionId, u64)>>>,
|
|
) -> Result<(SessionId, Option<i32>), ErrorDto> {
|
|
let sid = parse_session_id(&payload_string(&frame.payload, "sessionId")?)?;
|
|
ws_detach(frame, state, owned)?;
|
|
let output = state
|
|
.app
|
|
.close_terminal
|
|
.execute(CloseTerminalInput { session_id: sid })
|
|
.await
|
|
.map_err(ErrorDto::from)?;
|
|
Ok((sid, output.code))
|
|
}
|
|
|
|
async fn dispatch_http(
|
|
method: Method,
|
|
uri: Uri,
|
|
headers: HeaderMap,
|
|
body: Bytes,
|
|
state: Arc<ServerState>,
|
|
) -> Response<ResponseBody> {
|
|
if has_forbidden_query_secret(&uri) {
|
|
return error_response(
|
|
StatusCode::BAD_REQUEST,
|
|
"INVALID",
|
|
"secrets in URL query strings are forbidden",
|
|
None,
|
|
);
|
|
}
|
|
|
|
if !is_api_path(uri.path()) {
|
|
return serve_static(&state.config.web_root, method, uri.path()).await;
|
|
}
|
|
|
|
let origin = match validate_request_origin(&headers, &state.config) {
|
|
Ok(origin) => origin,
|
|
Err(response) => {
|
|
state.log_security(SecurityLogEvent::OriginRejected {
|
|
origin: request_origin(&headers),
|
|
route: uri.path().to_owned(),
|
|
});
|
|
return *response;
|
|
}
|
|
};
|
|
|
|
if method == Method::OPTIONS {
|
|
return cors_response(StatusCode::NO_CONTENT, origin.as_deref());
|
|
}
|
|
|
|
match (method, uri.path()) {
|
|
(Method::POST, "/api/pair") => pair(body, state, origin.as_deref()).await,
|
|
(Method::POST, "/api/logout") => {
|
|
let Some(token) = session_cookie(&headers) else {
|
|
return error_response(
|
|
StatusCode::UNAUTHORIZED,
|
|
"UNAUTHORIZED",
|
|
"missing session cookie",
|
|
origin.as_deref(),
|
|
);
|
|
};
|
|
if !state.revoke_session(&token) {
|
|
return error_response(
|
|
StatusCode::UNAUTHORIZED,
|
|
"UNAUTHORIZED",
|
|
"invalid session cookie",
|
|
origin.as_deref(),
|
|
);
|
|
}
|
|
state.log_security(SecurityLogEvent::SessionRevoked {
|
|
origin: origin.clone(),
|
|
});
|
|
logout_response(&state, origin.as_deref())
|
|
}
|
|
(Method::POST, "/api/invoke") => {
|
|
let Some(token) = session_cookie(&headers) else {
|
|
return error_response(
|
|
StatusCode::UNAUTHORIZED,
|
|
"UNAUTHORIZED",
|
|
"missing session cookie",
|
|
origin.as_deref(),
|
|
);
|
|
};
|
|
if !state.has_session(&token) {
|
|
return error_response(
|
|
StatusCode::UNAUTHORIZED,
|
|
"UNAUTHORIZED",
|
|
"invalid session cookie",
|
|
origin.as_deref(),
|
|
);
|
|
}
|
|
invoke(body, state, origin.as_deref()).await
|
|
}
|
|
(_, "/api/pair" | "/api/invoke" | "/api/logout") => error_response(
|
|
StatusCode::METHOD_NOT_ALLOWED,
|
|
"METHOD_NOT_ALLOWED",
|
|
"method not allowed",
|
|
origin.as_deref(),
|
|
),
|
|
_ if is_api_path(uri.path()) => error_response(
|
|
StatusCode::NOT_FOUND,
|
|
"NOT_FOUND",
|
|
"route not found",
|
|
origin.as_deref(),
|
|
),
|
|
_ => error_response(StatusCode::NOT_FOUND, "NOT_FOUND", "route not found", None),
|
|
}
|
|
}
|
|
|
|
fn is_api_path(path: &str) -> bool {
|
|
path == "/api" || path.starts_with("/api/")
|
|
}
|
|
|
|
async fn serve_static(web_root: &Path, method: Method, path: &str) -> Response<ResponseBody> {
|
|
if !matches!(method, Method::GET | Method::HEAD) {
|
|
return error_response(
|
|
StatusCode::METHOD_NOT_ALLOWED,
|
|
"METHOD_NOT_ALLOWED",
|
|
"method not allowed",
|
|
None,
|
|
);
|
|
}
|
|
|
|
let Some(route) = static_route(web_root, path) else {
|
|
return error_response(StatusCode::NOT_FOUND, "NOT_FOUND", "route not found", None);
|
|
};
|
|
|
|
let file = if route.candidate.is_file() {
|
|
route.candidate
|
|
} else if route.spa_fallback {
|
|
web_root.join("index.html")
|
|
} else {
|
|
return error_response(StatusCode::NOT_FOUND, "NOT_FOUND", "route not found", None);
|
|
};
|
|
|
|
let bytes = match tokio::fs::read(&file).await {
|
|
Ok(bytes) => bytes,
|
|
Err(_) => {
|
|
return error_response(StatusCode::NOT_FOUND, "NOT_FOUND", "route not found", None);
|
|
}
|
|
};
|
|
static_file_response(&file, bytes, method == Method::HEAD)
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
struct StaticRoute {
|
|
candidate: PathBuf,
|
|
spa_fallback: bool,
|
|
}
|
|
|
|
fn static_route(web_root: &Path, request_path: &str) -> Option<StaticRoute> {
|
|
if !request_path.starts_with('/') || request_path.contains('\\') {
|
|
return None;
|
|
}
|
|
let lower = request_path.to_ascii_lowercase();
|
|
if lower.contains("%2e") || lower.contains("%2f") || lower.contains("%5c") {
|
|
return None;
|
|
}
|
|
|
|
let mut candidate = web_root.to_path_buf();
|
|
let mut last = "";
|
|
for segment in request_path
|
|
.split('/')
|
|
.filter(|segment| !segment.is_empty())
|
|
{
|
|
if segment == "." || segment == ".." || segment.starts_with('.') || segment.contains('%') {
|
|
return None;
|
|
}
|
|
candidate.push(segment);
|
|
last = segment;
|
|
}
|
|
|
|
if request_path == "/" {
|
|
return Some(StaticRoute {
|
|
candidate: web_root.join("index.html"),
|
|
spa_fallback: false,
|
|
});
|
|
}
|
|
|
|
Some(StaticRoute {
|
|
candidate,
|
|
spa_fallback: !last.contains('.'),
|
|
})
|
|
}
|
|
|
|
fn static_file_response(path: &Path, bytes: Vec<u8>, head_only: bool) -> Response<ResponseBody> {
|
|
let is_index = path
|
|
.file_name()
|
|
.and_then(|name| name.to_str())
|
|
.is_some_and(|name| name == "index.html");
|
|
let body = if head_only { Vec::new() } else { bytes };
|
|
let mut response = Response::builder()
|
|
.status(StatusCode::OK)
|
|
.header(CONTENT_TYPE, content_type(path))
|
|
.header("x-content-type-options", "nosniff")
|
|
.body(Full::new(Bytes::from(body)))
|
|
.expect("static response is valid");
|
|
if is_index {
|
|
response.headers_mut().insert(
|
|
"content-security-policy",
|
|
HeaderValue::from_static(
|
|
"default-src 'self'; connect-src 'self' ws: wss:; img-src 'self' data: blob:; style-src 'self' 'unsafe-inline'; script-src 'self'; base-uri 'self'; frame-ancestors 'none'",
|
|
),
|
|
);
|
|
}
|
|
response
|
|
}
|
|
|
|
fn content_type(path: &Path) -> &'static str {
|
|
match path.extension().and_then(|ext| ext.to_str()).unwrap_or("") {
|
|
"html" => "text/html; charset=utf-8",
|
|
"js" | "mjs" => "text/javascript; charset=utf-8",
|
|
"css" => "text/css; charset=utf-8",
|
|
"json" => "application/json; charset=utf-8",
|
|
"svg" => "image/svg+xml",
|
|
"png" => "image/png",
|
|
"jpg" | "jpeg" => "image/jpeg",
|
|
"gif" => "image/gif",
|
|
"webp" => "image/webp",
|
|
"ico" => "image/x-icon",
|
|
"wasm" => "application/wasm",
|
|
"txt" => "text/plain; charset=utf-8",
|
|
_ => "application/octet-stream",
|
|
}
|
|
}
|
|
|
|
#[derive(Deserialize)]
|
|
struct PairRequest {
|
|
code: String,
|
|
}
|
|
|
|
async fn pair(
|
|
body: Bytes,
|
|
state: Arc<ServerState>,
|
|
origin: Option<&str>,
|
|
) -> Response<ResponseBody> {
|
|
let request = match serde_json::from_slice::<PairRequest>(&body) {
|
|
Ok(request) => request,
|
|
Err(err) => {
|
|
return error_response(
|
|
StatusCode::BAD_REQUEST,
|
|
"INVALID",
|
|
format!("invalid pairing request: {err}"),
|
|
origin,
|
|
);
|
|
}
|
|
};
|
|
|
|
if request.code != state.pairing_code() {
|
|
state.log_security(SecurityLogEvent::PairingFailed {
|
|
origin: origin.map(str::to_owned),
|
|
reason: "wrong_code",
|
|
});
|
|
return error_response(
|
|
StatusCode::FORBIDDEN,
|
|
"FORBIDDEN",
|
|
"invalid pairing code",
|
|
origin,
|
|
);
|
|
}
|
|
|
|
let token = state.create_session();
|
|
state.log_security(SecurityLogEvent::PairingSucceeded {
|
|
origin: origin.map(str::to_owned),
|
|
});
|
|
let cookie = Cookie::build((SESSION_COOKIE, token))
|
|
.path("/")
|
|
.http_only(true)
|
|
.secure(state.config.secure_cookie())
|
|
.same_site(SameSite::Strict)
|
|
.build();
|
|
let mut response = json_response(StatusCode::OK, &json!({ "paired": true }), origin);
|
|
response.headers_mut().insert(
|
|
SET_COOKIE,
|
|
HeaderValue::from_str(&cookie.to_string()).expect("session cookie is header-safe"),
|
|
);
|
|
response
|
|
}
|
|
|
|
fn logout_response(state: &ServerState, origin: Option<&str>) -> Response<ResponseBody> {
|
|
let mut response = json_response(StatusCode::OK, &json!({ "revoked": true }), origin);
|
|
let secure = if state.config.secure_cookie() {
|
|
"; Secure"
|
|
} else {
|
|
""
|
|
};
|
|
let cookie = format!("{SESSION_COOKIE}=; Path=/; HttpOnly{secure}; SameSite=Strict; Max-Age=0");
|
|
response.headers_mut().insert(
|
|
SET_COOKIE,
|
|
HeaderValue::from_str(&cookie).expect("session clearing cookie is header-safe"),
|
|
);
|
|
response
|
|
}
|
|
|
|
#[derive(Deserialize)]
|
|
struct InvokeRequest {
|
|
command: String,
|
|
#[serde(default)]
|
|
args: Value,
|
|
}
|
|
|
|
async fn invoke(
|
|
body: Bytes,
|
|
state: Arc<ServerState>,
|
|
origin: Option<&str>,
|
|
) -> Response<ResponseBody> {
|
|
let request = match serde_json::from_slice::<InvokeRequest>(&body) {
|
|
Ok(request) => request,
|
|
Err(err) => {
|
|
return error_response(
|
|
StatusCode::BAD_REQUEST,
|
|
"INVALID",
|
|
format!("invalid invoke request: {err}"),
|
|
origin,
|
|
);
|
|
}
|
|
};
|
|
|
|
let result = match request.command.as_str() {
|
|
"health" => invoke_health(&request.args, &state.app),
|
|
"list_projects" => invoke_list_projects(&state.app).await,
|
|
"open_project" => invoke_open_project(&request.args, &state.app).await,
|
|
"get_project_work_state" => invoke_get_project_work_state(&request.args, &state.app).await,
|
|
"list_background_tasks" => invoke_list_background_tasks(&request.args, &state.app).await,
|
|
"cancel_background_task" => invoke_cancel_background_task(&request.args, &state.app).await,
|
|
"retry_background_task" => invoke_retry_background_task(&request.args, &state.app).await,
|
|
_ => Err(ErrorDto {
|
|
code: "UNKNOWN_COMMAND".to_owned(),
|
|
message: format!("unknown command: {}", request.command),
|
|
}),
|
|
};
|
|
|
|
match result {
|
|
Ok(value) => json_response(StatusCode::OK, &value, origin),
|
|
Err(error) => error_dto_response(status_for_error(&error), error, origin),
|
|
}
|
|
}
|
|
|
|
fn invoke_health(args: &Value, state: &BackendCore) -> Result<Value, ErrorDto> {
|
|
let request = optional_request::<HealthRequestDto>(args)?;
|
|
let output = state
|
|
.health
|
|
.execute(request.unwrap_or_default().into())
|
|
.map(HealthResponseDto::from)
|
|
.map_err(ErrorDto::from)?;
|
|
serde_json::to_value(output).map_err(serialization_error)
|
|
}
|
|
|
|
async fn invoke_list_projects(state: &BackendCore) -> Result<Value, ErrorDto> {
|
|
let output = state
|
|
.list_projects
|
|
.execute()
|
|
.await
|
|
.map(ProjectListDto::from)
|
|
.map_err(ErrorDto::from)?;
|
|
serde_json::to_value(output).map_err(serialization_error)
|
|
}
|
|
|
|
async fn invoke_open_project(args: &Value, state: &BackendCore) -> Result<Value, ErrorDto> {
|
|
let project_id = args
|
|
.get("projectId")
|
|
.and_then(Value::as_str)
|
|
.ok_or_else(|| ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: "open_project requires args.projectId".to_owned(),
|
|
})?;
|
|
let project = resolve_project_readonly(project_id, state).await?;
|
|
serde_json::to_value(ProjectDto::from(project)).map_err(serialization_error)
|
|
}
|
|
|
|
async fn invoke_get_project_work_state(
|
|
args: &Value,
|
|
state: &BackendCore,
|
|
) -> Result<Value, ErrorDto> {
|
|
let project_id = args
|
|
.get("projectId")
|
|
.and_then(Value::as_str)
|
|
.ok_or_else(|| ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: "get_project_work_state requires args.projectId".to_owned(),
|
|
})?;
|
|
let project = resolve_project_readonly(project_id, state).await?;
|
|
let output = state
|
|
.get_project_work_state
|
|
.execute(GetProjectWorkStateInput { project })
|
|
.await
|
|
.map(ProjectWorkStateDto::from)
|
|
.map_err(ErrorDto::from)?;
|
|
serde_json::to_value(output).map_err(serialization_error)
|
|
}
|
|
|
|
async fn invoke_list_background_tasks(
|
|
args: &Value,
|
|
state: &BackendCore,
|
|
) -> Result<Value, ErrorDto> {
|
|
let project_id = args
|
|
.get("projectId")
|
|
.and_then(Value::as_str)
|
|
.ok_or_else(|| ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: "list_background_tasks requires args.projectId".to_owned(),
|
|
})
|
|
.and_then(parse_project_id)?;
|
|
let agent_id = args
|
|
.get("agentId")
|
|
.and_then(Value::as_str)
|
|
.map(parse_agent_id)
|
|
.transpose()?;
|
|
|
|
let mut by_id = std::collections::HashMap::<domain::TaskId, domain::BackgroundTask>::new();
|
|
if let Some(agent_id) = agent_id {
|
|
for task in state
|
|
.background_task_store
|
|
.list_open_for_agent(agent_id)
|
|
.await
|
|
.map_err(background_error)?
|
|
{
|
|
by_id.insert(task.id, task);
|
|
}
|
|
}
|
|
for task in state
|
|
.background_task_store
|
|
.list_undelivered_completions()
|
|
.await
|
|
.map_err(background_error)?
|
|
{
|
|
by_id.entry(task.id).or_insert(task);
|
|
}
|
|
|
|
let mut tasks = by_id
|
|
.into_values()
|
|
.filter(|task| task.project_id == project_id)
|
|
.filter(|task| agent_id.map_or(true, |agent_id| task.owner_agent_id == agent_id))
|
|
.collect::<Vec<_>>();
|
|
tasks.sort_by_key(|task| (task.created_at_ms, task.id));
|
|
let dto = tasks
|
|
.into_iter()
|
|
.map(BackgroundTaskDto::from)
|
|
.collect::<Vec<_>>();
|
|
serde_json::to_value(dto).map_err(serialization_error)
|
|
}
|
|
|
|
async fn invoke_cancel_background_task(
|
|
args: &Value,
|
|
state: &BackendCore,
|
|
) -> Result<Value, ErrorDto> {
|
|
let task_id = task_id_arg(args, "cancel_background_task")?;
|
|
verify_background_task_project(args, state, task_id).await?;
|
|
let output = state
|
|
.cancel_background_task
|
|
.execute(task_id)
|
|
.await
|
|
.map_err(ErrorDto::from)?;
|
|
serde_json::to_value(output.task.map(BackgroundTaskDto::from)).map_err(serialization_error)
|
|
}
|
|
|
|
async fn invoke_retry_background_task(
|
|
args: &Value,
|
|
state: &BackendCore,
|
|
) -> Result<Value, ErrorDto> {
|
|
let task_id = task_id_arg(args, "retry_background_task")?;
|
|
verify_background_task_project(args, state, task_id).await?;
|
|
let output = state
|
|
.retry_background_task
|
|
.execute(task_id)
|
|
.await
|
|
.map_err(ErrorDto::from)?;
|
|
serde_json::to_value(BackgroundTaskDto::from(output.task)).map_err(serialization_error)
|
|
}
|
|
|
|
fn task_id_arg(args: &Value, command: &str) -> Result<domain::TaskId, ErrorDto> {
|
|
args.get("taskId")
|
|
.and_then(Value::as_str)
|
|
.ok_or_else(|| ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: format!("{command} requires args.taskId"),
|
|
})
|
|
.and_then(parse_task_id)
|
|
}
|
|
|
|
async fn verify_background_task_project(
|
|
args: &Value,
|
|
state: &BackendCore,
|
|
task_id: domain::TaskId,
|
|
) -> Result<(), ErrorDto> {
|
|
let Some(project_id) = args.get("projectId").and_then(Value::as_str) else {
|
|
return Ok(());
|
|
};
|
|
let project_id = parse_project_id(project_id)?;
|
|
let Some(task) = state
|
|
.background_task_store
|
|
.get(task_id)
|
|
.await
|
|
.map_err(background_error)?
|
|
else {
|
|
return Err(ErrorDto {
|
|
code: "NOT_FOUND".to_owned(),
|
|
message: format!("background task not found: {task_id}"),
|
|
});
|
|
};
|
|
if task.project_id != project_id {
|
|
return Err(ErrorDto {
|
|
code: "FORBIDDEN".to_owned(),
|
|
message: "background task does not belong to requested project".to_owned(),
|
|
});
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
fn background_error(err: BackgroundTaskPortError) -> ErrorDto {
|
|
let code = match &err {
|
|
BackgroundTaskPortError::NotFound => "NOT_FOUND",
|
|
BackgroundTaskPortError::AlreadyExists | BackgroundTaskPortError::Invalid(_) => "INVALID",
|
|
BackgroundTaskPortError::Runner(_) => "PROCESS",
|
|
BackgroundTaskPortError::Store(_) => "STORE",
|
|
};
|
|
ErrorDto {
|
|
code: code.to_owned(),
|
|
message: err.to_string(),
|
|
}
|
|
}
|
|
|
|
async fn resolve_project_readonly(
|
|
project_id: &str,
|
|
state: &BackendCore,
|
|
) -> Result<Project, ErrorDto> {
|
|
let id = parse_project_id(project_id)?;
|
|
state
|
|
.open_project
|
|
.execute(OpenProjectInput { project_id: id })
|
|
.await
|
|
.map(|output| output.project)
|
|
.map_err(ErrorDto::from)
|
|
}
|
|
|
|
fn optional_request<T>(args: &Value) -> Result<Option<T>, ErrorDto>
|
|
where
|
|
T: for<'de> Deserialize<'de>,
|
|
{
|
|
match args {
|
|
Value::Object(map) => match map.get("request") {
|
|
Some(value) => serde_json::from_value(value.clone())
|
|
.map(Some)
|
|
.map_err(invalid_args_error),
|
|
None if map.is_empty() => Ok(None),
|
|
None => serde_json::from_value(args.clone())
|
|
.map(Some)
|
|
.map_err(invalid_args_error),
|
|
},
|
|
Value::Null => Ok(None),
|
|
_ => Err(ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: "args must be an object".to_owned(),
|
|
}),
|
|
}
|
|
}
|
|
|
|
fn validate_request_origin(
|
|
headers: &HeaderMap,
|
|
config: &ServerConfig,
|
|
) -> Result<Option<String>, Box<Response<ResponseBody>>> {
|
|
let Some(origin) = headers.get(ORIGIN).and_then(|value| value.to_str().ok()) else {
|
|
return Err(Box::new(error_response(
|
|
StatusCode::FORBIDDEN,
|
|
"FORBIDDEN",
|
|
"missing Origin header",
|
|
None,
|
|
)));
|
|
};
|
|
|
|
if origin_allowed(origin, config) {
|
|
Ok(Some(origin.to_owned()))
|
|
} else {
|
|
Err(Box::new(error_response(
|
|
StatusCode::FORBIDDEN,
|
|
"FORBIDDEN",
|
|
"origin not allowed",
|
|
None,
|
|
)))
|
|
}
|
|
}
|
|
|
|
fn request_origin(headers: &HeaderMap) -> Option<String> {
|
|
headers
|
|
.get(ORIGIN)
|
|
.and_then(|value| value.to_str().ok())
|
|
.map(str::to_owned)
|
|
}
|
|
|
|
fn origin_allowed(origin: &str, config: &ServerConfig) -> bool {
|
|
if let Some(public_origin) = &config.public_origin {
|
|
return origin == public_origin;
|
|
}
|
|
|
|
if !config.listen.ip().is_loopback() {
|
|
return false;
|
|
}
|
|
|
|
let port = config.listen.port();
|
|
origin == format!("http://127.0.0.1:{port}")
|
|
|| origin == format!("http://localhost:{port}")
|
|
|| origin == format!("http://[::1]:{port}")
|
|
}
|
|
|
|
fn has_forbidden_query_secret(uri: &Uri) -> bool {
|
|
uri.query().is_some_and(|query| {
|
|
query.split('&').any(|pair| {
|
|
let key = pair.split_once('=').map_or(pair, |(key, _)| key);
|
|
matches!(
|
|
key.to_ascii_lowercase().as_str(),
|
|
"token" | "secret" | "session" | "code"
|
|
)
|
|
})
|
|
})
|
|
}
|
|
|
|
fn session_cookie(headers: &HeaderMap) -> Option<String> {
|
|
headers
|
|
.get(COOKIE)
|
|
.and_then(|value| value.to_str().ok())
|
|
.and_then(|raw| {
|
|
raw.split(';').find_map(|part| {
|
|
let (name, value) = part.trim().split_once('=')?;
|
|
(name == SESSION_COOKIE).then(|| value.to_owned())
|
|
})
|
|
})
|
|
}
|
|
|
|
fn json_response(
|
|
status: StatusCode,
|
|
value: &Value,
|
|
origin: Option<&str>,
|
|
) -> Response<ResponseBody> {
|
|
let body = serde_json::to_vec(value).expect("JSON value serializes");
|
|
let mut response = Response::new(Full::new(Bytes::from(body)));
|
|
*response.status_mut() = status;
|
|
response
|
|
.headers_mut()
|
|
.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
|
|
add_cors_headers(&mut response, origin);
|
|
response
|
|
}
|
|
|
|
fn error_dto_response(
|
|
status: StatusCode,
|
|
error: ErrorDto,
|
|
origin: Option<&str>,
|
|
) -> Response<ResponseBody> {
|
|
let value = serde_json::to_value(error).expect("ErrorDto serializes");
|
|
json_response(status, &value, origin)
|
|
}
|
|
|
|
fn error_response(
|
|
status: StatusCode,
|
|
code: impl Into<String>,
|
|
message: impl Into<String>,
|
|
origin: Option<&str>,
|
|
) -> Response<ResponseBody> {
|
|
error_dto_response(
|
|
status,
|
|
ErrorDto {
|
|
code: code.into(),
|
|
message: message.into(),
|
|
},
|
|
origin,
|
|
)
|
|
}
|
|
|
|
fn cors_response(status: StatusCode, origin: Option<&str>) -> Response<ResponseBody> {
|
|
let mut response = Response::new(Full::new(Bytes::new()));
|
|
*response.status_mut() = status;
|
|
add_cors_headers(&mut response, origin);
|
|
response
|
|
}
|
|
|
|
fn add_cors_headers(response: &mut Response<ResponseBody>, origin: Option<&str>) {
|
|
if let Some(origin) = origin {
|
|
if let Ok(origin) = HeaderValue::from_str(origin) {
|
|
response
|
|
.headers_mut()
|
|
.insert("access-control-allow-origin", origin);
|
|
response.headers_mut().insert(
|
|
"access-control-allow-credentials",
|
|
HeaderValue::from_static("true"),
|
|
);
|
|
response.headers_mut().insert(
|
|
"access-control-allow-headers",
|
|
HeaderValue::from_static("content-type"),
|
|
);
|
|
response.headers_mut().insert(
|
|
"access-control-allow-methods",
|
|
HeaderValue::from_static("POST, OPTIONS"),
|
|
);
|
|
}
|
|
}
|
|
}
|
|
|
|
fn status_for_error(error: &ErrorDto) -> StatusCode {
|
|
match error.code.as_str() {
|
|
"UNKNOWN_COMMAND" => StatusCode::BAD_REQUEST,
|
|
"INVALID" => StatusCode::BAD_REQUEST,
|
|
"NOT_FOUND" => StatusCode::NOT_FOUND,
|
|
"FORBIDDEN" => StatusCode::FORBIDDEN,
|
|
"UNAUTHORIZED" => StatusCode::UNAUTHORIZED,
|
|
_ => StatusCode::INTERNAL_SERVER_ERROR,
|
|
}
|
|
}
|
|
|
|
fn validate_origin(origin: &str) -> Result<(), String> {
|
|
if origin == "*" || origin.contains('*') {
|
|
return Err("--public-origin must be exact, not a wildcard".to_owned());
|
|
}
|
|
if origin.ends_with('/') {
|
|
return Err("--public-origin must not end with '/'".to_owned());
|
|
}
|
|
if !(origin.starts_with("https://") || origin.starts_with("http://")) {
|
|
return Err("--public-origin must be an HTTP(S) origin".to_owned());
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
fn default_app_data_dir() -> PathBuf {
|
|
if let Some(path) = env::var_os("IDEA_APP_DATA_DIR") {
|
|
return PathBuf::from(path);
|
|
}
|
|
if let Some(path) = env::var_os("XDG_DATA_HOME") {
|
|
return PathBuf::from(path).join("app.idea.ide");
|
|
}
|
|
if let Some(home) = env::var_os("HOME") {
|
|
return PathBuf::from(home).join(".local/share/app.idea.ide");
|
|
}
|
|
env::current_dir()
|
|
.unwrap_or_else(|_| PathBuf::from("."))
|
|
.join(".ideai/app-data")
|
|
}
|
|
|
|
fn new_pairing_code() -> String {
|
|
Uuid::new_v4()
|
|
.simple()
|
|
.to_string()
|
|
.chars()
|
|
.take(8)
|
|
.collect::<String>()
|
|
.to_ascii_uppercase()
|
|
}
|
|
|
|
fn new_session_token() -> String {
|
|
format!("{}{}", Uuid::new_v4().simple(), Uuid::new_v4().simple())
|
|
}
|
|
|
|
fn invalid_args_error(err: serde_json::Error) -> ErrorDto {
|
|
ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: format!("invalid command args: {err}"),
|
|
}
|
|
}
|
|
|
|
fn serialization_error(err: serde_json::Error) -> ErrorDto {
|
|
ErrorDto {
|
|
code: "SERIALIZATION".to_owned(),
|
|
message: format!("failed to serialize response: {err}"),
|
|
}
|
|
}
|
|
|
|
#[derive(Debug, Clone, Deserialize, Serialize)]
|
|
struct ClientFrame {
|
|
id: String,
|
|
kind: String,
|
|
#[serde(default)]
|
|
payload: Value,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize)]
|
|
struct ServerFrame {
|
|
kind: String,
|
|
#[serde(skip_serializing_if = "Option::is_none", rename = "replyTo")]
|
|
reply_to: Option<String>,
|
|
payload: Value,
|
|
}
|
|
|
|
impl ServerFrame {
|
|
fn attached(
|
|
request_id: &str,
|
|
session: &TerminalSessionDto,
|
|
scrollback: Vec<u8>,
|
|
next_seq: u64,
|
|
gap: bool,
|
|
) -> Self {
|
|
let scrollback = if scrollback.is_empty() {
|
|
Vec::new()
|
|
} else {
|
|
vec![json!({
|
|
"seq": 0,
|
|
"bytesBase64": base64_encode(&scrollback),
|
|
})]
|
|
};
|
|
Self {
|
|
kind: "terminal.attached".to_owned(),
|
|
reply_to: Some(request_id.to_owned()),
|
|
payload: json!({
|
|
"session": {
|
|
"sessionId": session.session_id,
|
|
"rows": session.rows,
|
|
"cols": session.cols,
|
|
},
|
|
"scrollback": scrollback,
|
|
"nextSeq": next_seq,
|
|
"status": "running",
|
|
"gap": gap,
|
|
"assignedConversationId": session.assigned_conversation_id,
|
|
}),
|
|
}
|
|
}
|
|
|
|
fn output(session_id: SessionId, seq: u64, bytes: Vec<u8>) -> Self {
|
|
Self {
|
|
kind: "terminal.output".to_owned(),
|
|
reply_to: None,
|
|
payload: json!({
|
|
"sessionId": session_id.to_string(),
|
|
"seq": seq,
|
|
"bytesBase64": base64_encode(&bytes),
|
|
}),
|
|
}
|
|
}
|
|
|
|
fn status(session_id: SessionId, status: &str, exit_code: Option<i32>) -> Self {
|
|
Self {
|
|
kind: "terminal.status".to_owned(),
|
|
reply_to: None,
|
|
payload: json!({
|
|
"sessionId": session_id.to_string(),
|
|
"status": status,
|
|
"exitCode": exit_code,
|
|
}),
|
|
}
|
|
}
|
|
|
|
fn domain_event(event: &DomainEvent) -> Self {
|
|
Self {
|
|
kind: "event.domain".to_owned(),
|
|
reply_to: None,
|
|
payload: serde_json::to_value(DomainEventDto::from(event))
|
|
.expect("DomainEventDto serializes"),
|
|
}
|
|
}
|
|
|
|
fn error(
|
|
request_id: Option<String>,
|
|
session_id: Option<String>,
|
|
code: impl Into<String>,
|
|
message: impl Into<String>,
|
|
) -> Self {
|
|
Self {
|
|
kind: "error".to_owned(),
|
|
reply_to: request_id,
|
|
payload: json!({
|
|
"sessionId": session_id,
|
|
"code": code.into(),
|
|
"message": message.into(),
|
|
}),
|
|
}
|
|
}
|
|
|
|
fn pong() -> Self {
|
|
Self {
|
|
kind: "pong".to_owned(),
|
|
reply_to: None,
|
|
payload: json!({}),
|
|
}
|
|
}
|
|
}
|
|
|
|
struct WsPtySink {
|
|
session_id: SessionId,
|
|
tx: mpsc::Sender<ServerFrame>,
|
|
seq: AtomicU64,
|
|
}
|
|
|
|
impl WsPtySink {
|
|
fn new(session_id: SessionId, tx: mpsc::Sender<ServerFrame>, next_seq: u64) -> Self {
|
|
Self {
|
|
session_id,
|
|
tx,
|
|
seq: AtomicU64::new(next_seq),
|
|
}
|
|
}
|
|
}
|
|
|
|
impl OutputSink<PtyChunk> for WsPtySink {
|
|
fn send(&self, item: PtyChunk) -> Result<(), OutputSinkError> {
|
|
let seq = self.seq.fetch_add(1, Ordering::Relaxed);
|
|
self.tx
|
|
.try_send(ServerFrame::output(self.session_id, seq, item))
|
|
.map_err(|err| match err {
|
|
mpsc::error::TrySendError::Full(_) => OutputSinkError::Full,
|
|
mpsc::error::TrySendError::Closed(_) => OutputSinkError::Closed,
|
|
})
|
|
}
|
|
}
|
|
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
enum WsOpcode {
|
|
Text,
|
|
Binary,
|
|
Close,
|
|
Ping,
|
|
Pong,
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
struct WsFrame {
|
|
opcode: WsOpcode,
|
|
payload: Vec<u8>,
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
enum WsFrameError {
|
|
Incomplete,
|
|
Protocol(String),
|
|
TooLarge,
|
|
}
|
|
|
|
impl std::fmt::Display for WsFrameError {
|
|
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
|
match self {
|
|
Self::Incomplete => f.write_str("incomplete websocket frame"),
|
|
Self::Protocol(message) => write!(f, "websocket protocol error: {message}"),
|
|
Self::TooLarge => f.write_str("websocket frame payload too large"),
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn read_ws_frame<R>(reader: &mut R) -> Result<WsFrame, WsFrameError>
|
|
where
|
|
R: AsyncRead + Unpin,
|
|
{
|
|
let mut header = [0_u8; 2];
|
|
reader
|
|
.read_exact(&mut header)
|
|
.await
|
|
.map_err(|_| WsFrameError::Incomplete)?;
|
|
let len_code = header[1] & 0x7f;
|
|
let masked = header[1] & 0x80 != 0;
|
|
let mut rest = Vec::new();
|
|
let extra_len = match len_code {
|
|
126 => 2,
|
|
127 => 8,
|
|
_ => 0,
|
|
};
|
|
rest.resize(extra_len + usize::from(masked) * 4, 0);
|
|
if !rest.is_empty() {
|
|
reader
|
|
.read_exact(&mut rest)
|
|
.await
|
|
.map_err(|_| WsFrameError::Incomplete)?;
|
|
}
|
|
let mut raw = Vec::with_capacity(2 + rest.len());
|
|
raw.extend_from_slice(&header);
|
|
raw.extend_from_slice(&rest);
|
|
let payload_len = decoded_payload_len(&raw)?;
|
|
if payload_len > WS_MAX_PAYLOAD {
|
|
return Err(WsFrameError::TooLarge);
|
|
}
|
|
let mut payload = vec![0_u8; payload_len];
|
|
reader
|
|
.read_exact(&mut payload)
|
|
.await
|
|
.map_err(|_| WsFrameError::Incomplete)?;
|
|
raw.extend_from_slice(&payload);
|
|
decode_client_ws_frame(&raw)
|
|
}
|
|
|
|
fn decode_client_ws_frame(raw: &[u8]) -> Result<WsFrame, WsFrameError> {
|
|
if raw.len() < 2 {
|
|
return Err(WsFrameError::Incomplete);
|
|
}
|
|
let fin = raw[0] & 0x80 != 0;
|
|
if !fin {
|
|
return Err(WsFrameError::Protocol(
|
|
"fragmented frames are not supported".to_owned(),
|
|
));
|
|
}
|
|
let opcode = match raw[0] & 0x0f {
|
|
0x1 => WsOpcode::Text,
|
|
0x2 => WsOpcode::Binary,
|
|
0x8 => WsOpcode::Close,
|
|
0x9 => WsOpcode::Ping,
|
|
0xA => WsOpcode::Pong,
|
|
other => {
|
|
return Err(WsFrameError::Protocol(format!(
|
|
"unsupported opcode {other}"
|
|
)))
|
|
}
|
|
};
|
|
let masked = raw[1] & 0x80 != 0;
|
|
if !masked {
|
|
return Err(WsFrameError::Protocol(
|
|
"client frames must be masked".to_owned(),
|
|
));
|
|
}
|
|
let len_code = raw[1] & 0x7f;
|
|
let mut offset = 2;
|
|
let len = match len_code {
|
|
126 => {
|
|
if raw.len() < offset + 2 {
|
|
return Err(WsFrameError::Incomplete);
|
|
}
|
|
let len = u16::from_be_bytes([raw[offset], raw[offset + 1]]) as usize;
|
|
offset += 2;
|
|
len
|
|
}
|
|
127 => {
|
|
if raw.len() < offset + 8 {
|
|
return Err(WsFrameError::Incomplete);
|
|
}
|
|
let len = u64::from_be_bytes(raw[offset..offset + 8].try_into().expect("8 bytes"));
|
|
offset += 8;
|
|
usize::try_from(len).map_err(|_| WsFrameError::TooLarge)?
|
|
}
|
|
len => usize::from(len),
|
|
};
|
|
if len > WS_MAX_PAYLOAD {
|
|
return Err(WsFrameError::TooLarge);
|
|
}
|
|
if raw.len() < offset + 4 + len {
|
|
return Err(WsFrameError::Incomplete);
|
|
}
|
|
let mask = &raw[offset..offset + 4];
|
|
offset += 4;
|
|
let mut payload = raw[offset..offset + len].to_vec();
|
|
for (i, byte) in payload.iter_mut().enumerate() {
|
|
*byte ^= mask[i % 4];
|
|
}
|
|
Ok(WsFrame { opcode, payload })
|
|
}
|
|
|
|
fn decoded_payload_len(raw_header: &[u8]) -> Result<usize, WsFrameError> {
|
|
let len_code = raw_header[1] & 0x7f;
|
|
let offset = 2;
|
|
let len = match len_code {
|
|
126 => {
|
|
if raw_header.len() < offset + 2 {
|
|
return Err(WsFrameError::Incomplete);
|
|
}
|
|
u16::from_be_bytes([raw_header[offset], raw_header[offset + 1]]) as usize
|
|
}
|
|
127 => {
|
|
if raw_header.len() < offset + 8 {
|
|
return Err(WsFrameError::Incomplete);
|
|
}
|
|
let len =
|
|
u64::from_be_bytes(raw_header[offset..offset + 8].try_into().expect("8 bytes"));
|
|
usize::try_from(len).map_err(|_| WsFrameError::TooLarge)?
|
|
}
|
|
len => usize::from(len),
|
|
};
|
|
Ok(len)
|
|
}
|
|
|
|
fn encode_ws_frame(opcode: WsOpcode, payload: &[u8]) -> Vec<u8> {
|
|
let opcode = match opcode {
|
|
WsOpcode::Text => 0x1,
|
|
WsOpcode::Binary => 0x2,
|
|
WsOpcode::Close => 0x8,
|
|
WsOpcode::Ping => 0x9,
|
|
WsOpcode::Pong => 0xA,
|
|
};
|
|
let mut out = vec![0x80 | opcode];
|
|
match payload.len() {
|
|
0..=125 => out.push(payload.len() as u8),
|
|
126..=65535 => {
|
|
out.push(126);
|
|
out.extend_from_slice(&(payload.len() as u16).to_be_bytes());
|
|
}
|
|
len => {
|
|
out.push(127);
|
|
out.extend_from_slice(&(len as u64).to_be_bytes());
|
|
}
|
|
}
|
|
out.extend_from_slice(payload);
|
|
out
|
|
}
|
|
|
|
fn websocket_accept(key: &str) -> String {
|
|
let mut input = Vec::with_capacity(key.len() + WS_MAGIC.len());
|
|
input.extend_from_slice(key.as_bytes());
|
|
input.extend_from_slice(WS_MAGIC.as_bytes());
|
|
base64_encode(&sha1_digest(&input))
|
|
}
|
|
|
|
fn base64_encode(bytes: &[u8]) -> String {
|
|
use base64::Engine as _;
|
|
base64::engine::general_purpose::STANDARD.encode(bytes)
|
|
}
|
|
|
|
fn base64_decode(raw: &str) -> Result<Vec<u8>, ErrorDto> {
|
|
use base64::Engine as _;
|
|
base64::engine::general_purpose::STANDARD
|
|
.decode(raw)
|
|
.map_err(|err| ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: format!("invalid base64 bytes: {err}"),
|
|
})
|
|
}
|
|
|
|
fn sha1_digest(input: &[u8]) -> [u8; 20] {
|
|
let mut h0: u32 = 0x67452301;
|
|
let mut h1: u32 = 0xEFCDAB89;
|
|
let mut h2: u32 = 0x98BADCFE;
|
|
let mut h3: u32 = 0x10325476;
|
|
let mut h4: u32 = 0xC3D2E1F0;
|
|
|
|
let bit_len = (input.len() as u64) * 8;
|
|
let mut msg = input.to_vec();
|
|
msg.push(0x80);
|
|
while (msg.len() % 64) != 56 {
|
|
msg.push(0);
|
|
}
|
|
msg.extend_from_slice(&bit_len.to_be_bytes());
|
|
|
|
for chunk in msg.chunks(64) {
|
|
let mut w = [0_u32; 80];
|
|
for (i, word) in w.iter_mut().take(16).enumerate() {
|
|
let j = i * 4;
|
|
*word = u32::from_be_bytes([chunk[j], chunk[j + 1], chunk[j + 2], chunk[j + 3]]);
|
|
}
|
|
for i in 16..80 {
|
|
w[i] = (w[i - 3] ^ w[i - 8] ^ w[i - 14] ^ w[i - 16]).rotate_left(1);
|
|
}
|
|
let mut a = h0;
|
|
let mut b = h1;
|
|
let mut c = h2;
|
|
let mut d = h3;
|
|
let mut e = h4;
|
|
for (i, word) in w.iter().enumerate() {
|
|
let (f, k) = match i {
|
|
0..=19 => ((b & c) | ((!b) & d), 0x5A827999),
|
|
20..=39 => (b ^ c ^ d, 0x6ED9EBA1),
|
|
40..=59 => ((b & c) | (b & d) | (c & d), 0x8F1BBCDC),
|
|
_ => (b ^ c ^ d, 0xCA62C1D6),
|
|
};
|
|
let temp = a
|
|
.rotate_left(5)
|
|
.wrapping_add(f)
|
|
.wrapping_add(e)
|
|
.wrapping_add(k)
|
|
.wrapping_add(*word);
|
|
e = d;
|
|
d = c;
|
|
c = b.rotate_left(30);
|
|
b = a;
|
|
a = temp;
|
|
}
|
|
h0 = h0.wrapping_add(a);
|
|
h1 = h1.wrapping_add(b);
|
|
h2 = h2.wrapping_add(c);
|
|
h3 = h3.wrapping_add(d);
|
|
h4 = h4.wrapping_add(e);
|
|
}
|
|
|
|
let mut out = [0_u8; 20];
|
|
out[0..4].copy_from_slice(&h0.to_be_bytes());
|
|
out[4..8].copy_from_slice(&h1.to_be_bytes());
|
|
out[8..12].copy_from_slice(&h2.to_be_bytes());
|
|
out[12..16].copy_from_slice(&h3.to_be_bytes());
|
|
out[16..20].copy_from_slice(&h4.to_be_bytes());
|
|
out
|
|
}
|
|
|
|
fn payload_string(payload: &Value, key: &str) -> Result<String, ErrorDto> {
|
|
payload
|
|
.get(key)
|
|
.and_then(Value::as_str)
|
|
.map(ToOwned::to_owned)
|
|
.ok_or_else(|| ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: format!("payload requires {key}"),
|
|
})
|
|
}
|
|
|
|
fn payload_u16(payload: &Value, key: &str) -> Option<u16> {
|
|
payload
|
|
.get(key)
|
|
.and_then(Value::as_u64)
|
|
.and_then(|value| u16::try_from(value).ok())
|
|
}
|
|
|
|
fn payload_bytes(payload: &Value) -> Result<Vec<u8>, ErrorDto> {
|
|
if let Some(raw) = payload.get("bytesBase64").and_then(Value::as_str) {
|
|
base64_decode(raw)
|
|
} else if let Some(raw) = payload.get("bytes").and_then(Value::as_str) {
|
|
base64_decode(raw)
|
|
} else {
|
|
Err(ErrorDto {
|
|
code: "INVALID".to_owned(),
|
|
message: "payload requires bytesBase64".to_owned(),
|
|
})
|
|
}
|
|
}
|
|
|
|
fn payload_session_id(payload: &Value) -> Option<String> {
|
|
payload
|
|
.get("sessionId")
|
|
.and_then(Value::as_str)
|
|
.map(ToOwned::to_owned)
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use application::{
|
|
CreateAgentInput, CreateProjectInput, LaunchAgentOutput, SaveProfileInput,
|
|
StructuredSessionDescriptor,
|
|
};
|
|
use domain::ports::EventBus;
|
|
use domain::{
|
|
AgentId, AgentProfile, BackgroundTask, BackgroundTaskKind, BackgroundTaskWakePolicy,
|
|
ContextInjection, NodeId, ProfileId, ProjectId, ProjectPath, PtySize, SessionKind,
|
|
SessionStatus, SessionStrategy, TaskId, TerminalSession,
|
|
};
|
|
use http::header::HeaderName;
|
|
use std::time::Duration;
|
|
|
|
static ENV_LOCK: Mutex<()> = Mutex::new(());
|
|
|
|
struct EnvVarGuard {
|
|
saved: Vec<(&'static str, Option<std::ffi::OsString>)>,
|
|
}
|
|
|
|
impl EnvVarGuard {
|
|
fn new(keys: &[&'static str]) -> Self {
|
|
Self {
|
|
saved: keys
|
|
.iter()
|
|
.map(|key| (*key, std::env::var_os(key)))
|
|
.collect(),
|
|
}
|
|
}
|
|
}
|
|
|
|
impl Drop for EnvVarGuard {
|
|
fn drop(&mut self) {
|
|
for (key, value) in self.saved.iter().rev() {
|
|
if let Some(value) = value {
|
|
std::env::set_var(key, value);
|
|
} else {
|
|
std::env::remove_var(key);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
#[derive(Default)]
|
|
struct RecordingSecurityLogger {
|
|
events: Mutex<Vec<SecurityLogEvent>>,
|
|
}
|
|
|
|
impl RecordingSecurityLogger {
|
|
fn events(&self) -> Vec<SecurityLogEvent> {
|
|
self.events.lock().unwrap().clone()
|
|
}
|
|
}
|
|
|
|
impl SecurityLogger for RecordingSecurityLogger {
|
|
fn log(&self, event: SecurityLogEvent) {
|
|
self.events.lock().unwrap().push(event);
|
|
}
|
|
}
|
|
|
|
fn create_web_root() -> PathBuf {
|
|
let root = std::env::temp_dir().join(format!("idea-web-root-{}", Uuid::new_v4()));
|
|
std::fs::create_dir_all(root.join("assets")).unwrap();
|
|
std::fs::write(
|
|
root.join("index.html"),
|
|
r#"<!doctype html><div id="root"></div><script type="module" src="/assets/app.js"></script>"#,
|
|
)
|
|
.unwrap();
|
|
std::fs::write(root.join("assets").join("app.js"), "console.log('idea');").unwrap();
|
|
std::fs::write(root.join("assets").join("style.css"), "body{margin:0}").unwrap();
|
|
root
|
|
}
|
|
|
|
fn test_config() -> ServerConfig {
|
|
ServerConfig {
|
|
listen: "127.0.0.1:17373".parse().unwrap(),
|
|
public_origin: None,
|
|
allow_remote: false,
|
|
trust_reverse_proxy: false,
|
|
app_data_dir: std::env::temp_dir().join(format!("idea-server-test-{}", Uuid::new_v4())),
|
|
web_root: create_web_root(),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn default_app_data_dir_uses_tauri_identifier_with_env_precedence() {
|
|
let _lock = ENV_LOCK.lock().unwrap();
|
|
let _guard = EnvVarGuard::new(&["IDEA_APP_DATA_DIR", "XDG_DATA_HOME", "HOME"]);
|
|
let root = std::env::temp_dir().join(format!("idea-app-data-env-{}", Uuid::new_v4()));
|
|
let home = root.join("home");
|
|
let xdg = root.join("xdg");
|
|
let override_dir = root.join("override");
|
|
|
|
std::env::remove_var("IDEA_APP_DATA_DIR");
|
|
std::env::remove_var("XDG_DATA_HOME");
|
|
std::env::set_var("HOME", &home);
|
|
assert_eq!(
|
|
default_app_data_dir(),
|
|
home.join(".local/share/app.idea.ide")
|
|
);
|
|
|
|
std::env::set_var("XDG_DATA_HOME", &xdg);
|
|
assert_eq!(default_app_data_dir(), xdg.join("app.idea.ide"));
|
|
|
|
std::env::set_var("IDEA_APP_DATA_DIR", &override_dir);
|
|
assert_eq!(default_app_data_dir(), override_dir);
|
|
}
|
|
|
|
fn state() -> Arc<ServerState> {
|
|
Arc::new(ServerState::new_for_test(test_config(), "PAIR1234"))
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn run_embedded_stop_shuts_down_accept_loop() {
|
|
let config = ServerConfig {
|
|
listen: "127.0.0.1:0".parse().unwrap(),
|
|
..test_config()
|
|
};
|
|
|
|
let handle = run_embedded(config).await.unwrap();
|
|
|
|
assert_eq!(handle.state(), EmbeddedServerState::Running);
|
|
assert_ne!(handle.url(), "http://127.0.0.1:0");
|
|
handle.stop().await.unwrap();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn run_embedded_with_core_uses_injected_core_for_http_invokes() {
|
|
let core = Arc::new(BackendCore::build(
|
|
std::env::temp_dir().join(format!("idea-shared-core-{}", Uuid::new_v4())),
|
|
));
|
|
let project_id = create_project_on_core(&core, "Shared Core").await;
|
|
let config = ServerConfig {
|
|
listen: "127.0.0.1:0".parse().unwrap(),
|
|
app_data_dir: std::env::temp_dir()
|
|
.join(format!("idea-unused-embedded-core-{}", Uuid::new_v4())),
|
|
..test_config()
|
|
};
|
|
|
|
let handle = run_embedded_with_core(config, Arc::clone(&core))
|
|
.await
|
|
.unwrap();
|
|
let (pair_status, _, pair_headers) = embedded_json_request(
|
|
handle.url(),
|
|
"/api/pair",
|
|
json!({ "code": handle.pairing_code() }),
|
|
&[],
|
|
)
|
|
.await;
|
|
let cookie = pair_headers
|
|
.get(SET_COOKIE)
|
|
.unwrap()
|
|
.to_str()
|
|
.unwrap()
|
|
.split(';')
|
|
.next()
|
|
.unwrap()
|
|
.to_owned();
|
|
let (invoke_status, body, _) = embedded_json_request(
|
|
handle.url(),
|
|
"/api/invoke",
|
|
json!({ "command": "list_projects", "args": {} }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
handle.stop().await.unwrap();
|
|
|
|
assert_eq!(pair_status, StatusCode::OK);
|
|
assert_eq!(invoke_status, StatusCode::OK);
|
|
let projects = body
|
|
.as_array()
|
|
.expect("ProjectListDto is a transparent array");
|
|
let project = projects
|
|
.iter()
|
|
.find(|project| project["id"] == project_id)
|
|
.expect("project created on injected core is listed");
|
|
assert_eq!(project["name"], "Shared Core");
|
|
}
|
|
|
|
async fn request(
|
|
state: Arc<ServerState>,
|
|
method: Method,
|
|
uri: &str,
|
|
body: Value,
|
|
extra_headers: &[(&str, &str)],
|
|
) -> Response<ResponseBody> {
|
|
let mut builder = Request::builder()
|
|
.method(method)
|
|
.uri(uri)
|
|
.header(ORIGIN, "http://127.0.0.1:17373")
|
|
.header(CONTENT_TYPE, "application/json");
|
|
for (name, value) in extra_headers {
|
|
builder = builder.header(HeaderName::from_bytes(name.as_bytes()).unwrap(), *value);
|
|
}
|
|
handle_request(
|
|
builder
|
|
.body(Full::new(Bytes::from(serde_json::to_vec(&body).unwrap())))
|
|
.unwrap(),
|
|
state,
|
|
)
|
|
.await
|
|
}
|
|
|
|
async fn response_json(response: Response<ResponseBody>) -> (StatusCode, Value, HeaderMap) {
|
|
let status = response.status();
|
|
let headers = response.headers().clone();
|
|
let body = response.into_body().collect().await.unwrap().to_bytes();
|
|
let value = if body.is_empty() {
|
|
Value::Null
|
|
} else {
|
|
serde_json::from_slice(&body).unwrap()
|
|
};
|
|
(status, value, headers)
|
|
}
|
|
|
|
async fn embedded_json_request(
|
|
base_url: &str,
|
|
path: &str,
|
|
body: Value,
|
|
extra_headers: &[(&str, &str)],
|
|
) -> (StatusCode, Value, HeaderMap) {
|
|
let addr = base_url
|
|
.strip_prefix("http://")
|
|
.expect("embedded server URL is HTTP");
|
|
let body = serde_json::to_vec(&body).unwrap();
|
|
let mut request = format!(
|
|
"POST {path} HTTP/1.1\r\n\
|
|
host: {addr}\r\n\
|
|
origin: {base_url}\r\n\
|
|
content-type: application/json\r\n\
|
|
content-length: {}\r\n",
|
|
body.len()
|
|
)
|
|
.into_bytes();
|
|
for (name, value) in extra_headers {
|
|
request.extend_from_slice(name.as_bytes());
|
|
request.extend_from_slice(b": ");
|
|
request.extend_from_slice(value.as_bytes());
|
|
request.extend_from_slice(b"\r\n");
|
|
}
|
|
request.extend_from_slice(b"\r\n");
|
|
request.extend_from_slice(&body);
|
|
|
|
let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
|
|
stream.write_all(&request).await.unwrap();
|
|
let mut response = Vec::new();
|
|
stream.read_to_end(&mut response).await.unwrap();
|
|
|
|
let header_end = find_header_end(&response).expect("response has headers");
|
|
let head = std::str::from_utf8(&response[..header_end]).unwrap();
|
|
let mut lines = head.split("\r\n");
|
|
let status_line = lines.next().unwrap();
|
|
let status_code = status_line
|
|
.split_whitespace()
|
|
.nth(1)
|
|
.unwrap()
|
|
.parse::<u16>()
|
|
.unwrap();
|
|
let mut headers = HeaderMap::new();
|
|
for line in lines {
|
|
let Some((name, value)) = line.split_once(':') else {
|
|
continue;
|
|
};
|
|
headers.insert(
|
|
HeaderName::from_bytes(name.trim().as_bytes()).unwrap(),
|
|
HeaderValue::from_str(value.trim()).unwrap(),
|
|
);
|
|
}
|
|
let body = &response[header_end + 4..];
|
|
let value = if body.is_empty() {
|
|
Value::Null
|
|
} else {
|
|
serde_json::from_slice(body).unwrap()
|
|
};
|
|
(StatusCode::from_u16(status_code).unwrap(), value, headers)
|
|
}
|
|
|
|
async fn response_bytes(response: Response<ResponseBody>) -> (StatusCode, Bytes, HeaderMap) {
|
|
let status = response.status();
|
|
let headers = response.headers().clone();
|
|
let body = response.into_body().collect().await.unwrap().to_bytes();
|
|
(status, body, headers)
|
|
}
|
|
|
|
async fn pair_and_cookie(state: Arc<ServerState>) -> String {
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/pair",
|
|
json!({ "code": "PAIR1234" }),
|
|
&[],
|
|
)
|
|
.await;
|
|
response
|
|
.headers()
|
|
.get(SET_COOKIE)
|
|
.unwrap()
|
|
.to_str()
|
|
.unwrap()
|
|
.split(';')
|
|
.next()
|
|
.unwrap()
|
|
.to_owned()
|
|
}
|
|
|
|
async fn create_project_for_test(state: &Arc<ServerState>, name: &str) -> String {
|
|
let root = std::env::temp_dir()
|
|
.join(format!("idea-server-project-{}", Uuid::new_v4()))
|
|
.to_string_lossy()
|
|
.into_owned();
|
|
let output = state
|
|
.app
|
|
.create_project
|
|
.execute(CreateProjectInput {
|
|
name: name.to_owned(),
|
|
root,
|
|
remote: None,
|
|
default_profile_id: None,
|
|
})
|
|
.await
|
|
.expect("test project is created");
|
|
output.project.id.to_string()
|
|
}
|
|
|
|
async fn create_project_on_core(core: &BackendCore, name: &str) -> String {
|
|
let root = std::env::temp_dir()
|
|
.join(format!(
|
|
"idea-server-shared-core-project-{}",
|
|
Uuid::new_v4()
|
|
))
|
|
.to_string_lossy()
|
|
.into_owned();
|
|
let output = core
|
|
.create_project
|
|
.execute(CreateProjectInput {
|
|
name: name.to_owned(),
|
|
root,
|
|
remote: None,
|
|
default_profile_id: None,
|
|
})
|
|
.await
|
|
.expect("test project is created on injected core");
|
|
output.project.id.to_string()
|
|
}
|
|
|
|
async fn create_background_task_for_test(
|
|
state: &Arc<ServerState>,
|
|
project_id: &str,
|
|
label: &str,
|
|
) -> (TaskId, AgentId) {
|
|
let task_id = TaskId::from_uuid(Uuid::new_v4());
|
|
let owner = AgentId::from_uuid(Uuid::new_v4());
|
|
let project_id = ProjectId::from_uuid(Uuid::parse_str(project_id).unwrap());
|
|
let task = BackgroundTask::new(
|
|
task_id,
|
|
project_id,
|
|
owner,
|
|
BackgroundTaskKind::Command {
|
|
label: label.to_owned(),
|
|
},
|
|
BackgroundTaskWakePolicy::WakeOwner,
|
|
1_000,
|
|
None,
|
|
)
|
|
.expect("test background task is valid");
|
|
state
|
|
.app
|
|
.background_task_store
|
|
.create(&task)
|
|
.await
|
|
.expect("test background task is stored");
|
|
(task_id, owner)
|
|
}
|
|
|
|
async fn create_raw_cli_agent_for_test(
|
|
state: &Arc<ServerState>,
|
|
name: &str,
|
|
) -> (String, String) {
|
|
let project_id = create_project_for_test(state, name).await;
|
|
let project = resolve_project_readonly(&project_id, &state.app)
|
|
.await
|
|
.expect("test project resolves");
|
|
let profile_id = ProfileId::from_uuid(Uuid::new_v4());
|
|
let profile = AgentProfile::new(
|
|
profile_id,
|
|
format!("{name} shell profile"),
|
|
"/bin/sh",
|
|
vec![
|
|
"-c".to_owned(),
|
|
"printf agent-ready; sleep 30".to_owned(),
|
|
"idea-test-sh".to_owned(),
|
|
],
|
|
ContextInjection::env("IDEA_CONTEXT").expect("valid env injection"),
|
|
None,
|
|
"{agentRunDir}",
|
|
Some(
|
|
SessionStrategy::new(Some("--session-id".to_owned()), "--resume")
|
|
.expect("valid session strategy"),
|
|
),
|
|
)
|
|
.expect("valid raw CLI profile");
|
|
state
|
|
.app
|
|
.save_profile
|
|
.execute(SaveProfileInput { profile })
|
|
.await
|
|
.expect("test profile saved");
|
|
let agent = state
|
|
.app
|
|
.create_agent
|
|
.execute(CreateAgentInput {
|
|
project,
|
|
name: format!("{name} agent"),
|
|
profile_id,
|
|
initial_content: Some("Test agent context".to_owned()),
|
|
})
|
|
.await
|
|
.expect("test agent created")
|
|
.agent;
|
|
(project_id, agent.id.to_string())
|
|
}
|
|
|
|
fn ws_headers(origin: &str, cookie: Option<&str>) -> HeaderMap {
|
|
let mut headers = HeaderMap::new();
|
|
headers.insert(ORIGIN, HeaderValue::from_str(origin).unwrap());
|
|
headers.insert("upgrade", HeaderValue::from_static("websocket"));
|
|
headers.insert("connection", HeaderValue::from_static("Upgrade"));
|
|
headers.insert(
|
|
"sec-websocket-key",
|
|
HeaderValue::from_static("dGhlIHNhbXBsZSBub25jZQ=="),
|
|
);
|
|
if let Some(cookie) = cookie {
|
|
headers.insert(COOKIE, HeaderValue::from_str(cookie).unwrap());
|
|
}
|
|
headers
|
|
}
|
|
|
|
fn masked_text_frame(text: &str) -> Vec<u8> {
|
|
let payload = text.as_bytes();
|
|
let mask = [1_u8, 2, 3, 4];
|
|
let mut out = vec![0x81];
|
|
out.push(0x80 | payload.len() as u8);
|
|
out.extend_from_slice(&mask);
|
|
for (i, byte) in payload.iter().enumerate() {
|
|
out.push(byte ^ mask[i % 4]);
|
|
}
|
|
out
|
|
}
|
|
|
|
async fn recv_server_frame(rx: &mut mpsc::Receiver<ServerFrame>) -> ServerFrame {
|
|
tokio::time::timeout(Duration::from_secs(2), rx.recv())
|
|
.await
|
|
.expect("server frame received before timeout")
|
|
.expect("server frame channel is open")
|
|
}
|
|
|
|
async fn recv_server_frame_where(
|
|
rx: &mut mpsc::Receiver<ServerFrame>,
|
|
mut predicate: impl FnMut(&ServerFrame) -> bool,
|
|
) -> ServerFrame {
|
|
for _ in 0..16 {
|
|
let frame = recv_server_frame(rx).await;
|
|
if predicate(&frame) {
|
|
return frame;
|
|
}
|
|
}
|
|
panic!("matching server frame was not received");
|
|
}
|
|
|
|
fn drain_server_frames(rx: &mut mpsc::Receiver<ServerFrame>) {
|
|
while rx.try_recv().is_ok() {}
|
|
}
|
|
|
|
fn ws_frame(id: &str, kind: &str, payload: Value) -> ClientFrame {
|
|
ClientFrame {
|
|
id: id.to_owned(),
|
|
kind: kind.to_owned(),
|
|
payload,
|
|
}
|
|
}
|
|
|
|
fn attached_session_id(frame: &ServerFrame) -> String {
|
|
assert_eq!(frame.kind, "terminal.attached");
|
|
frame.payload["session"]["sessionId"]
|
|
.as_str()
|
|
.expect("attached frame carries sessionId")
|
|
.to_owned()
|
|
}
|
|
|
|
fn open_terminal_payload(command: &str, args: Vec<&str>) -> Value {
|
|
json!({
|
|
"request": {
|
|
"cwd": std::env::temp_dir().to_string_lossy(),
|
|
"rows": 24,
|
|
"cols": 80,
|
|
"command": command,
|
|
"args": args,
|
|
}
|
|
})
|
|
}
|
|
|
|
fn agent_launch_payload(
|
|
project_id: &str,
|
|
agent_id: &str,
|
|
node_id: &str,
|
|
conversation_id: Option<&str>,
|
|
) -> Value {
|
|
json!({
|
|
"projectId": project_id,
|
|
"agentId": agent_id,
|
|
"nodeId": node_id,
|
|
"rows": 24,
|
|
"cols": 80,
|
|
"conversationId": conversation_id,
|
|
})
|
|
}
|
|
|
|
async fn close_ws_session(state: &Arc<ServerState>, session_id: &str) {
|
|
let (tx, mut rx) = mpsc::channel(8);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"close-cleanup",
|
|
"terminal.close",
|
|
json!({ "sessionId": session_id }),
|
|
),
|
|
state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let _ = recv_server_frame(&mut rx).await;
|
|
}
|
|
|
|
async fn wait_for_scrollback(
|
|
state: &Arc<ServerState>,
|
|
session_id: &str,
|
|
needle: &[u8],
|
|
) -> Vec<u8> {
|
|
let sid = parse_session_id(session_id).expect("test session id parses");
|
|
let handle = PtyHandle { session_id: sid };
|
|
for _ in 0..100 {
|
|
if let Ok(scrollback) = state.app.pty_port.scrollback(&handle) {
|
|
if scrollback
|
|
.windows(needle.len())
|
|
.any(|window| window == needle)
|
|
{
|
|
return scrollback;
|
|
}
|
|
}
|
|
tokio::time::sleep(Duration::from_millis(20)).await;
|
|
}
|
|
panic!(
|
|
"scrollback did not contain {:?}",
|
|
String::from_utf8_lossy(needle)
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_accept_matches_rfc_example() {
|
|
assert_eq!(
|
|
websocket_accept("dGhlIHNhbXBsZSBub25jZQ=="),
|
|
"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_upgrade_requires_valid_cookie() {
|
|
let state = state();
|
|
let headers = ws_headers("http://127.0.0.1:17373", None);
|
|
|
|
let response = validate_ws_upgrade(&headers, &state).unwrap_err();
|
|
|
|
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_upgrade_rejects_invalid_cookie() {
|
|
let state = state();
|
|
let headers = ws_headers(
|
|
"http://127.0.0.1:17373",
|
|
Some("idea_session=not-a-valid-session"),
|
|
);
|
|
|
|
let response = validate_ws_upgrade(&headers, &state).unwrap_err();
|
|
|
|
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_upgrade_requires_allowed_origin() {
|
|
let state = state();
|
|
let token = state.create_session();
|
|
let headers = ws_headers(
|
|
"https://evil.example",
|
|
Some(&format!("idea_session={token}")),
|
|
);
|
|
|
|
let response = validate_ws_upgrade(&headers, &state).unwrap_err();
|
|
|
|
assert_eq!(response.status(), StatusCode::FORBIDDEN);
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_upgrade_accepts_valid_cookie_and_origin() {
|
|
let state = state();
|
|
let token = state.create_session();
|
|
let headers = ws_headers(
|
|
"http://127.0.0.1:17373",
|
|
Some(&format!("idea_session={token}")),
|
|
);
|
|
|
|
let accept = validate_ws_upgrade(&headers, &state).unwrap();
|
|
|
|
assert_eq!(accept, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=");
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_client_text_frame_decodes_masked_json() {
|
|
let raw = masked_text_frame(r#"{"kind":"ping"}"#);
|
|
|
|
let frame = decode_client_ws_frame(&raw).unwrap();
|
|
|
|
assert_eq!(frame.opcode, WsOpcode::Text);
|
|
assert_eq!(frame.payload, br#"{"kind":"ping"}"#);
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_client_json_frame_round_trips_base64_bytes() {
|
|
let sid = domain::SessionId::new_random().to_string();
|
|
let value = json!({
|
|
"id": "req-1",
|
|
"kind": "input",
|
|
"payload": {
|
|
"sessionId": sid,
|
|
"bytesBase64": "AQID"
|
|
}
|
|
});
|
|
|
|
let frame: ClientFrame = serde_json::from_value(value.clone()).unwrap();
|
|
let roundtrip = serde_json::to_value(&frame).unwrap();
|
|
|
|
assert_eq!(roundtrip, value);
|
|
assert_eq!(payload_bytes(&frame.payload).unwrap(), vec![1, 2, 3]);
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_server_attached_frame_uses_contract_shape() {
|
|
let sid = domain::SessionId::new_random().to_string();
|
|
let frame = ServerFrame::attached(
|
|
"req-attach",
|
|
&TerminalSessionDto {
|
|
session_id: sid.clone(),
|
|
cwd: "/tmp".to_owned(),
|
|
rows: 24,
|
|
cols: 80,
|
|
assigned_conversation_id: None,
|
|
engine_session_id: None,
|
|
cell_kind: backend::dto::CellKind::Pty,
|
|
},
|
|
vec![4, 5, 6],
|
|
1,
|
|
true,
|
|
);
|
|
let value = serde_json::to_value(frame).unwrap();
|
|
|
|
assert_eq!(value["kind"], "terminal.attached");
|
|
assert_eq!(value["replyTo"], "req-attach");
|
|
assert_eq!(value["payload"]["session"]["sessionId"], sid);
|
|
assert_eq!(value["payload"]["scrollback"][0]["seq"], 0);
|
|
assert_eq!(value["payload"]["scrollback"][0]["bytesBase64"], "BAUG");
|
|
assert_eq!(value["payload"]["nextSeq"], 1);
|
|
assert_eq!(value["payload"]["gap"], true);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_app_ping_emits_pong() {
|
|
let state = state();
|
|
let (tx, mut rx) = mpsc::channel(4);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
|
|
handle_client_frame(ws_frame("ping-1", "ping", json!({})), &state, &tx, &owned).await;
|
|
let frame = recv_server_frame(&mut rx).await;
|
|
|
|
assert_eq!(frame.kind, "pong");
|
|
assert!(frame.payload.as_object().unwrap().is_empty());
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_domain_event_frame_uses_contract_shape() {
|
|
let project_id = ProjectId::from_uuid(Uuid::new_v4());
|
|
let frame = ServerFrame::domain_event(&DomainEvent::ProjectCreated { project_id });
|
|
|
|
assert_eq!(frame.kind, "event.domain");
|
|
assert!(frame.reply_to.is_none());
|
|
assert_eq!(frame.payload["type"], "projectCreated");
|
|
assert_eq!(frame.payload["projectId"], project_id.to_string());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_domain_event_relay_forwards_domain_events() {
|
|
let state = state();
|
|
let project_id = ProjectId::from_uuid(Uuid::new_v4());
|
|
let event_rx = state.app.event_bus.raw_receiver();
|
|
let (tx, mut rx) = mpsc::channel(4);
|
|
let relay = spawn_ws_domain_event_relay(event_rx, tx);
|
|
|
|
state
|
|
.app
|
|
.event_bus
|
|
.publish(DomainEvent::ProjectCreated { project_id });
|
|
let frame = tokio::time::timeout(Duration::from_secs(1), rx.recv())
|
|
.await
|
|
.expect("event frame is forwarded")
|
|
.expect("event frame channel is open");
|
|
relay.abort();
|
|
|
|
assert_eq!(frame.kind, "event.domain");
|
|
assert_eq!(frame.payload["type"], "projectCreated");
|
|
assert_eq!(frame.payload["projectId"], project_id.to_string());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_domain_event_relay_forwards_background_completion() {
|
|
let state = state();
|
|
let project_id = ProjectId::from_uuid(Uuid::new_v4());
|
|
let task_id = TaskId::from_uuid(Uuid::new_v4());
|
|
let owner = AgentId::from_uuid(Uuid::new_v4());
|
|
let event_rx = state.app.event_bus.raw_receiver();
|
|
let (tx, mut rx) = mpsc::channel(4);
|
|
let relay = spawn_ws_domain_event_relay(event_rx, tx);
|
|
|
|
state
|
|
.app
|
|
.event_bus
|
|
.publish(DomainEvent::BackgroundTaskCompleted {
|
|
project_id,
|
|
task_id,
|
|
owner_agent_id: owner,
|
|
});
|
|
let frame = tokio::time::timeout(Duration::from_secs(1), rx.recv())
|
|
.await
|
|
.expect("background completion event is forwarded")
|
|
.expect("event frame channel is open");
|
|
relay.abort();
|
|
|
|
assert_eq!(frame.kind, "event.domain");
|
|
assert_eq!(frame.payload["type"], "backgroundTaskChanged");
|
|
assert_eq!(frame.payload["projectId"], project_id.to_string());
|
|
assert_eq!(frame.payload["taskId"], task_id.to_string());
|
|
assert_eq!(frame.payload["agentId"], owner.to_string());
|
|
assert_eq!(frame.payload["state"], "completed");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_domain_event_relay_excludes_pty_output() {
|
|
let state = state();
|
|
let project_id = ProjectId::from_uuid(Uuid::new_v4());
|
|
let event_rx = state.app.event_bus.raw_receiver();
|
|
let (tx, mut rx) = mpsc::channel(4);
|
|
let relay = spawn_ws_domain_event_relay(event_rx, tx);
|
|
|
|
state.app.event_bus.publish(DomainEvent::PtyOutput {
|
|
session_id: SessionId::new_random(),
|
|
bytes: b"hidden from global live bus".to_vec(),
|
|
});
|
|
state
|
|
.app
|
|
.event_bus
|
|
.publish(DomainEvent::ProjectCreated { project_id });
|
|
let frame = tokio::time::timeout(Duration::from_secs(1), rx.recv())
|
|
.await
|
|
.expect("sentinel event is forwarded")
|
|
.expect("event frame channel is open");
|
|
relay.abort();
|
|
|
|
assert_eq!(frame.kind, "event.domain");
|
|
assert_eq!(frame.payload["type"], "projectCreated");
|
|
assert_eq!(frame.payload["projectId"], project_id.to_string());
|
|
assert!(
|
|
rx.try_recv().is_err(),
|
|
"PtyOutput must not be emitted on the global event.domain stream"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_domain_event_relay_drops_when_client_queue_is_full() {
|
|
let state = state();
|
|
let event_rx = state.app.event_bus.raw_receiver();
|
|
let (tx, mut rx) = mpsc::channel(1);
|
|
assert!(tx.try_send(ServerFrame::pong()).is_ok());
|
|
let relay = spawn_ws_domain_event_relay(event_rx, tx);
|
|
|
|
state.app.event_bus.publish(DomainEvent::ProjectCreated {
|
|
project_id: ProjectId::from_uuid(Uuid::new_v4()),
|
|
});
|
|
tokio::time::sleep(Duration::from_millis(20)).await;
|
|
let first = rx.try_recv().expect("pre-filled frame remains queued");
|
|
relay.abort();
|
|
|
|
assert_eq!(first.kind, "pong");
|
|
assert!(
|
|
rx.try_recv().is_err(),
|
|
"full client queue should drop the live event instead of buffering"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_open_terminal_emits_attached_with_empty_scrollback() {
|
|
let state = state();
|
|
let (tx, mut rx) = mpsc::channel(16);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"open-1",
|
|
"terminal.open",
|
|
open_terminal_payload("/bin/sh", vec!["-c", "sleep 30"]),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let frame = recv_server_frame(&mut rx).await;
|
|
let session_id = attached_session_id(&frame);
|
|
|
|
assert_eq!(frame.reply_to.as_deref(), Some("open-1"));
|
|
assert_eq!(frame.payload["scrollback"].as_array().unwrap().len(), 0);
|
|
assert_eq!(frame.payload["nextSeq"], 0);
|
|
|
|
close_ws_session(&state, &session_id).await;
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_attach_terminal_replays_scrollback_in_attached_ack() {
|
|
let state = state();
|
|
let (tx, mut rx) = mpsc::channel(32);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"open-replay",
|
|
"terminal.open",
|
|
open_terminal_payload("/bin/sh", vec!["-c", "printf replay-ready; sleep 30"]),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let opened = recv_server_frame(&mut rx).await;
|
|
let session_id = attached_session_id(&opened);
|
|
|
|
let scrollback = wait_for_scrollback(&state, &session_id, b"replay-ready").await;
|
|
assert!(scrollback.len() <= 100 * 1024);
|
|
drain_server_frames(&mut rx);
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"attach-replay",
|
|
"terminal.attach",
|
|
json!({
|
|
"sessionId": session_id,
|
|
"lastSeq": 0,
|
|
"rows": 30,
|
|
"cols": 100,
|
|
}),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let attached = recv_server_frame(&mut rx).await;
|
|
let replay = attached.payload["scrollback"].as_array().unwrap();
|
|
|
|
assert_eq!(attached.kind, "terminal.attached");
|
|
assert_eq!(attached.reply_to.as_deref(), Some("attach-replay"));
|
|
assert_eq!(attached.payload["session"]["rows"], 30);
|
|
assert_eq!(attached.payload["session"]["cols"], 100);
|
|
assert_eq!(attached.payload["gap"], true);
|
|
assert_eq!(attached.payload["nextSeq"], 1);
|
|
assert_eq!(
|
|
base64_decode(replay[0]["bytesBase64"].as_str().unwrap()).unwrap(),
|
|
scrollback
|
|
);
|
|
|
|
close_ws_session(
|
|
&state,
|
|
attached.payload["session"]["sessionId"].as_str().unwrap(),
|
|
)
|
|
.await;
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_close_terminal_emits_exited_status_and_releases_session() {
|
|
let state = state();
|
|
let (tx, mut rx) = mpsc::channel(16);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"open-close",
|
|
"terminal.open",
|
|
open_terminal_payload("/bin/sh", vec!["-c", "sleep 30"]),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let opened = recv_server_frame(&mut rx).await;
|
|
let session_id = attached_session_id(&opened);
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"close-1",
|
|
"terminal.close",
|
|
json!({ "sessionId": session_id }),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let status = recv_server_frame(&mut rx).await;
|
|
|
|
assert_eq!(status.kind, "terminal.status");
|
|
assert_eq!(status.payload["status"], "exited");
|
|
assert_eq!(status.payload["sessionId"], session_id);
|
|
assert_eq!(state.ws_pty_bridge.active_sessions(), 0);
|
|
assert!(state.app.terminal_sessions.is_empty());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_launch_agent_emits_attached_with_assigned_conversation_id() {
|
|
let state = state();
|
|
let (project_id, agent_id) = create_raw_cli_agent_for_test(&state, "ws-launch").await;
|
|
let node_id = Uuid::new_v4().to_string();
|
|
let (tx, mut rx) = mpsc::channel(32);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"agent-launch-1",
|
|
"agent.launch",
|
|
agent_launch_payload(&project_id, &agent_id, &node_id, None),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let attached = recv_server_frame(&mut rx).await;
|
|
let session_id = attached_session_id(&attached);
|
|
|
|
assert_eq!(attached.kind, "terminal.attached");
|
|
assert_eq!(attached.reply_to.as_deref(), Some("agent-launch-1"));
|
|
assert_eq!(attached.payload["scrollback"].as_array().unwrap().len(), 0);
|
|
assert!(attached.payload["assignedConversationId"]
|
|
.as_str()
|
|
.and_then(|raw| Uuid::parse_str(raw).ok())
|
|
.is_some());
|
|
|
|
close_ws_session(&state, &session_id).await;
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_launch_agent_structured_is_unsupported() {
|
|
let state = state();
|
|
let (tx, mut rx) = mpsc::channel(4);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
let session_id = SessionId::new_random();
|
|
let agent_id = AgentId::from_uuid(Uuid::new_v4());
|
|
let node_id = NodeId::from_uuid(Uuid::new_v4());
|
|
let size = PtySize::new(24, 80).unwrap();
|
|
let mut session = TerminalSession::starting(
|
|
session_id,
|
|
node_id,
|
|
ProjectPath::new("/".to_owned()).unwrap(),
|
|
SessionKind::Agent { agent_id },
|
|
size,
|
|
);
|
|
session.status = SessionStatus::Running;
|
|
let frame = ws_frame("agent-structured", "agent.launch", json!({}));
|
|
let output = LaunchAgentOutput {
|
|
session,
|
|
assigned_conversation_id: Some("pair-conversation".to_owned()),
|
|
engine_session_id: Some("engine-session".to_owned()),
|
|
structured: Some(StructuredSessionDescriptor {
|
|
session_id,
|
|
agent_id,
|
|
node_id,
|
|
conversation_id: Some("engine-session".to_owned()),
|
|
}),
|
|
profile: None,
|
|
};
|
|
|
|
if let Err(err) = send_launch_agent_attached(&frame, &state, &tx, &owned, output).await {
|
|
tx.send(ServerFrame::error(
|
|
Some(frame.id),
|
|
payload_session_id(&frame.payload),
|
|
err.code,
|
|
err.message,
|
|
))
|
|
.await
|
|
.unwrap();
|
|
}
|
|
let error = recv_server_frame(&mut rx).await;
|
|
|
|
assert_eq!(error.kind, "error");
|
|
assert_eq!(error.reply_to.as_deref(), Some("agent-structured"));
|
|
assert_eq!(error.payload["code"], "UNSUPPORTED");
|
|
assert_eq!(
|
|
error.payload["message"],
|
|
"structured agent sessions do not stream over the PTY websocket"
|
|
);
|
|
assert_eq!(state.ws_pty_bridge.active_sessions(), 0);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_launch_agent_reattach_replays_scrollback_without_respawn() {
|
|
let state = state();
|
|
let (project_id, agent_id) = create_raw_cli_agent_for_test(&state, "ws-reattach").await;
|
|
let node_id = Uuid::new_v4().to_string();
|
|
let (tx, mut rx) = mpsc::channel(32);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"agent-launch-reattach",
|
|
"agent.launch",
|
|
agent_launch_payload(&project_id, &agent_id, &node_id, None),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let launched = recv_server_frame(&mut rx).await;
|
|
let session_id = attached_session_id(&launched);
|
|
let scrollback = wait_for_scrollback(&state, &session_id, b"agent-ready").await;
|
|
drain_server_frames(&mut rx);
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"agent-attach",
|
|
"terminal.attach",
|
|
json!({ "sessionId": session_id, "lastSeq": 0, "rows": 24, "cols": 80 }),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let attached = recv_server_frame_where(&mut rx, |frame| {
|
|
frame.kind == "terminal.attached" && frame.reply_to.as_deref() == Some("agent-attach")
|
|
})
|
|
.await;
|
|
let replay = attached.payload["scrollback"].as_array().unwrap();
|
|
|
|
assert_eq!(attached.kind, "terminal.attached");
|
|
assert_eq!(attached.payload["session"]["sessionId"], session_id);
|
|
assert_eq!(
|
|
base64_decode(replay[0]["bytesBase64"].as_str().unwrap()).unwrap(),
|
|
scrollback
|
|
);
|
|
let aid = parse_agent_id(&agent_id).unwrap();
|
|
assert_eq!(
|
|
state.app.terminal_sessions.sessions_for_agent(&aid).len(),
|
|
1
|
|
);
|
|
|
|
close_ws_session(&state, &session_id).await;
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_launch_agent_same_cell_is_idempotent_singleton() {
|
|
let state = state();
|
|
let (project_id, agent_id) = create_raw_cli_agent_for_test(&state, "ws-singleton").await;
|
|
let node_id = Uuid::new_v4().to_string();
|
|
let (tx, mut rx) = mpsc::channel(32);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"agent-launch-first",
|
|
"agent.launch",
|
|
agent_launch_payload(&project_id, &agent_id, &node_id, None),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let first = recv_server_frame(&mut rx).await;
|
|
let first_session = attached_session_id(&first);
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"agent-launch-second",
|
|
"agent.launch",
|
|
agent_launch_payload(
|
|
&project_id,
|
|
&agent_id,
|
|
&node_id,
|
|
first.payload["assignedConversationId"].as_str(),
|
|
),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let second = recv_server_frame_where(&mut rx, |frame| {
|
|
frame.kind == "terminal.attached"
|
|
&& frame.reply_to.as_deref() == Some("agent-launch-second")
|
|
})
|
|
.await;
|
|
|
|
assert_eq!(second.kind, "terminal.attached");
|
|
assert_eq!(attached_session_id(&second), first_session);
|
|
let aid = parse_agent_id(&agent_id).unwrap();
|
|
assert_eq!(
|
|
state.app.terminal_sessions.sessions_for_agent(&aid).len(),
|
|
1
|
|
);
|
|
|
|
close_ws_session(&state, &first_session).await;
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_launch_agent_different_cell_is_refused_by_singleton_guard() {
|
|
let state = state();
|
|
let (project_id, agent_id) =
|
|
create_raw_cli_agent_for_test(&state, "ws-singleton-refuse").await;
|
|
let first_node = Uuid::new_v4().to_string();
|
|
let second_node = Uuid::new_v4().to_string();
|
|
let (tx, mut rx) = mpsc::channel(32);
|
|
let owned = Arc::new(Mutex::new(Vec::new()));
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"agent-launch-first-cell",
|
|
"agent.launch",
|
|
agent_launch_payload(&project_id, &agent_id, &first_node, None),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let first = recv_server_frame(&mut rx).await;
|
|
let session_id = attached_session_id(&first);
|
|
|
|
handle_client_frame(
|
|
ws_frame(
|
|
"agent-launch-other-cell",
|
|
"agent.launch",
|
|
agent_launch_payload(&project_id, &agent_id, &second_node, None),
|
|
),
|
|
&state,
|
|
&tx,
|
|
&owned,
|
|
)
|
|
.await;
|
|
let error = recv_server_frame_where(&mut rx, |frame| {
|
|
frame.kind == "error" && frame.reply_to.as_deref() == Some("agent-launch-other-cell")
|
|
})
|
|
.await;
|
|
|
|
assert_eq!(error.kind, "error");
|
|
assert_eq!(error.reply_to.as_deref(), Some("agent-launch-other-cell"));
|
|
assert_eq!(error.payload["code"], "AGENT_ALREADY_RUNNING");
|
|
let aid = parse_agent_id(&agent_id).unwrap();
|
|
assert_eq!(
|
|
state.app.terminal_sessions.sessions_for_agent(&aid).len(),
|
|
1
|
|
);
|
|
|
|
close_ws_session(&state, &session_id).await;
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_client_unmasked_frame_is_rejected() {
|
|
let raw = [0x81, 0x02, b'h', b'i'];
|
|
|
|
let err = decode_client_ws_frame(&raw).unwrap_err();
|
|
|
|
assert!(matches!(err, WsFrameError::Protocol(_)));
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_payload_too_large_is_rejected() {
|
|
let len = (WS_MAX_PAYLOAD + 1) as u64;
|
|
let mut raw = vec![0x81, 0x80 | 127];
|
|
raw.extend_from_slice(&len.to_be_bytes());
|
|
raw.extend_from_slice(&[1, 2, 3, 4]);
|
|
|
|
let err = decode_client_ws_frame(&raw).unwrap_err();
|
|
|
|
assert_eq!(err, WsFrameError::TooLarge);
|
|
}
|
|
|
|
#[test]
|
|
fn websocket_server_frames_are_unmasked() {
|
|
let raw = encode_ws_frame(WsOpcode::Text, b"{}");
|
|
|
|
assert_eq!(raw[0], 0x81);
|
|
assert_eq!(raw[1] & 0x80, 0, "server frames must not be masked");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_sink_full_reports_backpressure_without_blocking() {
|
|
let (tx, mut rx) = mpsc::channel(1);
|
|
let sid = domain::SessionId::new_random();
|
|
let sink = WsPtySink::new(sid, tx, 0);
|
|
|
|
assert!(sink.send(vec![1]).is_ok());
|
|
assert_eq!(sink.send(vec![2]), Err(OutputSinkError::Full));
|
|
let frame = rx.recv().await.unwrap();
|
|
assert_eq!(frame.kind, "terminal.output");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn websocket_output_bridge_replaces_attachment_without_double_delivery() {
|
|
let bridge = OutputBridge::<SessionId, PtyChunk>::new();
|
|
let sid = domain::SessionId::new_random();
|
|
let (old_tx, mut old_rx) = mpsc::channel(4);
|
|
let (new_tx, mut new_rx) = mpsc::channel(4);
|
|
|
|
let old_gen = bridge.register(sid, Arc::new(WsPtySink::new(sid, old_tx, 0)));
|
|
let _new_gen = bridge.register(sid, Arc::new(WsPtySink::new(sid, new_tx, 0)));
|
|
|
|
assert!(bridge.send_output(&sid, vec![9]));
|
|
assert!(old_rx.try_recv().is_err());
|
|
let frame = new_rx.recv().await.unwrap();
|
|
assert_eq!(frame.kind, "terminal.output");
|
|
assert_eq!(frame.payload["bytesBase64"], "CQ==");
|
|
|
|
bridge.unregister_if(&sid, old_gen);
|
|
assert_eq!(bridge.active_sessions(), 1);
|
|
}
|
|
|
|
#[test]
|
|
fn public_bind_requires_explicit_remote_security() {
|
|
let config = ServerConfig {
|
|
listen: "0.0.0.0:17373".parse().unwrap(),
|
|
..test_config()
|
|
};
|
|
|
|
assert!(config.validate().is_err());
|
|
|
|
let config = ServerConfig {
|
|
allow_remote: true,
|
|
public_origin: Some("https://idea.example.com".to_owned()),
|
|
trust_reverse_proxy: true,
|
|
..config
|
|
};
|
|
assert!(config.validate().is_ok());
|
|
}
|
|
|
|
#[test]
|
|
fn allow_remote_requires_https_origin_and_reverse_proxy() {
|
|
let base = ServerConfig {
|
|
listen: "0.0.0.0:17373".parse().unwrap(),
|
|
allow_remote: true,
|
|
..test_config()
|
|
};
|
|
|
|
assert!(base.validate().is_err());
|
|
|
|
let with_http = ServerConfig {
|
|
public_origin: Some("http://idea.example.com".to_owned()),
|
|
trust_reverse_proxy: true,
|
|
..base.clone()
|
|
};
|
|
assert!(with_http.validate().is_err());
|
|
|
|
let without_proxy = ServerConfig {
|
|
public_origin: Some("https://idea.example.com".to_owned()),
|
|
trust_reverse_proxy: false,
|
|
..base
|
|
};
|
|
assert!(without_proxy.validate().is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn invoke_requires_valid_cookie() {
|
|
let state = state();
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "health", "args": {} }),
|
|
&[],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::UNAUTHORIZED);
|
|
assert_eq!(body["code"], "UNAUTHORIZED");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn invoke_rejects_invalid_session_cookie() {
|
|
let state = state();
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "health", "args": {} }),
|
|
&[("cookie", "idea_session=not-a-valid-session")],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::UNAUTHORIZED);
|
|
assert_eq!(body["code"], "UNAUTHORIZED");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn logout_revokes_session_for_http_and_websocket() {
|
|
let state = state();
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let response = request(
|
|
Arc::clone(&state),
|
|
Method::POST,
|
|
"/api/logout",
|
|
json!({}),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, headers) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert_eq!(body, json!({ "revoked": true }));
|
|
assert!(headers
|
|
.get(SET_COOKIE)
|
|
.unwrap()
|
|
.to_str()
|
|
.unwrap()
|
|
.contains("Max-Age=0"));
|
|
|
|
let response = request(
|
|
Arc::clone(&state),
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "health", "args": {} }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
assert_eq!(status, StatusCode::UNAUTHORIZED);
|
|
assert_eq!(body["code"], "UNAUTHORIZED");
|
|
|
|
let headers = ws_headers("http://127.0.0.1:17373", Some(&cookie));
|
|
let response = validate_ws_upgrade(&headers, &state).unwrap_err();
|
|
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn security_logs_pairing_origin_ws_and_revoke_without_secrets() {
|
|
let logger = Arc::new(RecordingSecurityLogger::default());
|
|
let state = Arc::new(ServerState::new_for_test_with_logger(
|
|
test_config(),
|
|
"PAIR1234",
|
|
logger.clone(),
|
|
));
|
|
|
|
let bad_origin = handle_request(
|
|
Request::builder()
|
|
.method(Method::POST)
|
|
.uri("/api/pair")
|
|
.header(ORIGIN, "https://evil.example")
|
|
.header(CONTENT_TYPE, "application/json")
|
|
.body(Full::new(Bytes::from_static(br#"{"code":"PAIR1234"}"#)))
|
|
.unwrap(),
|
|
Arc::clone(&state),
|
|
)
|
|
.await;
|
|
assert_eq!(bad_origin.status(), StatusCode::FORBIDDEN);
|
|
|
|
let wrong_code = request(
|
|
Arc::clone(&state),
|
|
Method::POST,
|
|
"/api/pair",
|
|
json!({ "code": "WRONG999" }),
|
|
&[],
|
|
)
|
|
.await;
|
|
assert_eq!(wrong_code.status(), StatusCode::FORBIDDEN);
|
|
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
let _ = request(
|
|
Arc::clone(&state),
|
|
Method::POST,
|
|
"/api/logout",
|
|
json!({}),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
|
|
let headers = ws_headers(
|
|
"http://127.0.0.1:17373",
|
|
Some("idea_session=invalid-token-value"),
|
|
);
|
|
let _ = validate_ws_upgrade(&headers, &state);
|
|
|
|
let events = logger.events();
|
|
assert!(events.iter().any(|event| matches!(
|
|
event,
|
|
SecurityLogEvent::OriginRejected { route, .. } if route == "/api/pair"
|
|
)));
|
|
assert!(events.iter().any(|event| matches!(
|
|
event,
|
|
SecurityLogEvent::PairingFailed {
|
|
reason: "wrong_code",
|
|
..
|
|
}
|
|
)));
|
|
assert!(events
|
|
.iter()
|
|
.any(|event| matches!(event, SecurityLogEvent::PairingSucceeded { .. })));
|
|
assert!(events
|
|
.iter()
|
|
.any(|event| matches!(event, SecurityLogEvent::SessionRevoked { .. })));
|
|
assert!(events.iter().any(|event| matches!(
|
|
event,
|
|
SecurityLogEvent::WsUpgradeRejected {
|
|
reason: "invalid_session",
|
|
..
|
|
}
|
|
)));
|
|
|
|
let lines = events.iter().map(security_log_line).collect::<Vec<_>>();
|
|
assert!(
|
|
!lines.iter().any(|line| line.contains("PAIR1234")
|
|
|| line.contains("WRONG999")
|
|
|| line.contains("invalid-token-value")),
|
|
"security logs must not leak pairing codes or session tokens"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn pairing_rejects_wrong_code() {
|
|
let state = state();
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/pair",
|
|
json!({ "code": "WRONG999" }),
|
|
&[],
|
|
)
|
|
.await;
|
|
let (status, body, headers) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::FORBIDDEN);
|
|
assert_eq!(body["code"], "FORBIDDEN");
|
|
assert!(
|
|
!headers.contains_key(SET_COOKIE),
|
|
"wrong pairing code must not issue a session cookie"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn pairing_sets_http_only_strict_cookie_for_loopback_dev() {
|
|
let state = state();
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/pair",
|
|
json!({ "code": "PAIR1234" }),
|
|
&[],
|
|
)
|
|
.await;
|
|
let (status, body, headers) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert_eq!(body, json!({ "paired": true }));
|
|
let cookie = headers.get(SET_COOKIE).unwrap().to_str().unwrap();
|
|
assert!(cookie.contains("idea_session="));
|
|
assert!(cookie.contains("HttpOnly"));
|
|
assert!(cookie.contains("SameSite=Strict"));
|
|
assert!(
|
|
!cookie.contains("Secure"),
|
|
"loopback dev over HTTP deliberately avoids Secure so browsers store it"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn pairing_sets_secure_cookie_for_remote_https_origin() {
|
|
let config = ServerConfig {
|
|
listen: "0.0.0.0:17373".parse().unwrap(),
|
|
public_origin: Some("https://idea.example.com".to_owned()),
|
|
allow_remote: true,
|
|
trust_reverse_proxy: true,
|
|
..test_config()
|
|
};
|
|
let state = Arc::new(ServerState::new_for_test(config, "PAIR1234"));
|
|
|
|
let mut req = Request::builder()
|
|
.method(Method::POST)
|
|
.uri("/api/pair")
|
|
.header(ORIGIN, "https://idea.example.com")
|
|
.header(CONTENT_TYPE, "application/json")
|
|
.body(Full::new(Bytes::from_static(br#"{"code":"PAIR1234"}"#)))
|
|
.unwrap();
|
|
*req.headers_mut().get_mut(CONTENT_TYPE).unwrap() =
|
|
HeaderValue::from_static("application/json");
|
|
|
|
let response = handle_request(req, state).await;
|
|
let (status, _, headers) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
let cookie = headers.get(SET_COOKIE).unwrap().to_str().unwrap();
|
|
assert!(cookie.contains("Secure"));
|
|
assert!(cookie.contains("HttpOnly"));
|
|
assert!(cookie.contains("SameSite=Strict"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn token_or_secret_in_query_string_is_refused() {
|
|
let state = state();
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke?token=abc",
|
|
json!({ "command": "health", "args": {} }),
|
|
&[],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::BAD_REQUEST);
|
|
assert_eq!(body["code"], "INVALID");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn serves_web_index_same_origin_with_csp() {
|
|
let state = state();
|
|
|
|
let response = request(state, Method::GET, "/", json!({}), &[]).await;
|
|
let (status, body, headers) = response_bytes(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert!(std::str::from_utf8(&body).unwrap().contains("id=\"root\""));
|
|
assert_eq!(
|
|
headers.get(CONTENT_TYPE).unwrap(),
|
|
"text/html; charset=utf-8"
|
|
);
|
|
assert!(headers.contains_key("content-security-policy"));
|
|
assert_eq!(headers.get("x-content-type-options").unwrap(), "nosniff");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn serves_static_assets_with_mime_types() {
|
|
let state = state();
|
|
|
|
let response = request(state, Method::GET, "/assets/app.js", json!({}), &[]).await;
|
|
let (status, body, headers) = response_bytes(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert_eq!(std::str::from_utf8(&body).unwrap(), "console.log('idea');");
|
|
assert_eq!(
|
|
headers.get(CONTENT_TYPE).unwrap(),
|
|
"text/javascript; charset=utf-8"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn serves_spa_fallback_for_non_api_routes_without_extension() {
|
|
let state = state();
|
|
|
|
let response = request(state, Method::GET, "/projects/abc", json!({}), &[]).await;
|
|
let (status, body, headers) = response_bytes(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert!(std::str::from_utf8(&body).unwrap().contains("id=\"root\""));
|
|
assert!(headers.contains_key("content-security-policy"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn missing_asset_and_api_routes_do_not_fallback_to_spa() {
|
|
let state = state();
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let missing_asset = request(
|
|
Arc::clone(&state),
|
|
Method::GET,
|
|
"/assets/missing.js",
|
|
json!({}),
|
|
&[],
|
|
)
|
|
.await;
|
|
let (asset_status, asset_body, _) = response_json(missing_asset).await;
|
|
assert_eq!(asset_status, StatusCode::NOT_FOUND);
|
|
assert_eq!(asset_body["code"], "NOT_FOUND");
|
|
|
|
let api_unknown = request(
|
|
state,
|
|
Method::GET,
|
|
"/api/unknown",
|
|
json!({}),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (api_status, api_body, _) = response_json(api_unknown).await;
|
|
assert_eq!(api_status, StatusCode::NOT_FOUND);
|
|
assert_eq!(api_body["code"], "NOT_FOUND");
|
|
}
|
|
|
|
#[test]
|
|
fn static_route_rejects_traversal_and_dotfiles() {
|
|
let root = PathBuf::from("/tmp/web");
|
|
|
|
assert!(static_route(&root, "/../secret").is_none());
|
|
assert!(static_route(&root, "/%2e%2e/secret").is_none());
|
|
assert!(static_route(&root, "/.env").is_none());
|
|
assert_eq!(
|
|
static_route(&root, "/dashboard").unwrap(),
|
|
StaticRoute {
|
|
candidate: root.join("dashboard"),
|
|
spa_fallback: true
|
|
}
|
|
);
|
|
assert_eq!(
|
|
static_route(&root, "/assets/app.js").unwrap(),
|
|
StaticRoute {
|
|
candidate: root.join("assets").join("app.js"),
|
|
spa_fallback: false
|
|
}
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn explicit_web_root_requires_index_html() {
|
|
let missing = std::env::temp_dir().join(format!("idea-empty-web-{}", Uuid::new_v4()));
|
|
std::fs::create_dir_all(&missing).unwrap();
|
|
|
|
assert!(validate_web_root(missing, "--web-root").is_err());
|
|
assert!(validate_web_root(create_web_root(), "--web-root").is_ok());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn non_allowed_origin_is_refused() {
|
|
let state = state();
|
|
let mut req = Request::builder()
|
|
.method(Method::POST)
|
|
.uri("/api/pair")
|
|
.header(ORIGIN, "https://evil.example")
|
|
.header(CONTENT_TYPE, "application/json")
|
|
.body(Full::new(Bytes::from_static(br#"{"code":"PAIR1234"}"#)))
|
|
.unwrap();
|
|
*req.headers_mut().get_mut(CONTENT_TYPE).unwrap() =
|
|
HeaderValue::from_static("application/json");
|
|
|
|
let response = handle_request(req, state).await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::FORBIDDEN);
|
|
assert_eq!(body["code"], "FORBIDDEN");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn unknown_allowlisted_command_returns_unknown_command() {
|
|
let state = state();
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "debug_dump", "args": {} }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::BAD_REQUEST);
|
|
assert_eq!(body["code"], "UNKNOWN_COMMAND");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn authorized_health_invoke_returns_health_report() {
|
|
let state = state();
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "health", "args": { "request": { "note": "hi" } } }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert_eq!(body["alive"], true);
|
|
assert_eq!(body["note"], "hi");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn authorized_list_projects_returns_tauri_project_list_contract() {
|
|
let state = state();
|
|
let project_id = create_project_for_test(&state, "Web Read").await;
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "list_projects", "args": {} }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
let projects = body
|
|
.as_array()
|
|
.expect("ProjectListDto is a transparent array");
|
|
let project = projects
|
|
.iter()
|
|
.find(|project| project["id"] == project_id)
|
|
.expect("created project is listed");
|
|
assert_eq!(project["name"], "Web Read");
|
|
assert!(project["root"]
|
|
.as_str()
|
|
.is_some_and(|root| root.contains("idea-server-project-")));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn authorized_open_project_returns_readonly_project_dto() {
|
|
let state = state();
|
|
let project_id = create_project_for_test(&state, "Readonly Open").await;
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "open_project", "args": { "projectId": project_id } }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert_eq!(body["id"], project_id);
|
|
assert_eq!(body["name"], "Readonly Open");
|
|
assert!(body["root"].is_string());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn authorized_get_project_work_state_returns_tauri_contract() {
|
|
let state = state();
|
|
let project_id = create_project_for_test(&state, "Readonly Work").await;
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "get_project_work_state", "args": { "projectId": project_id } }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert!(body["agents"].as_array().is_some());
|
|
assert!(body["conversations"].as_array().is_some());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn authorized_list_background_tasks_returns_tauri_contract() {
|
|
let state = state();
|
|
let project_id = create_project_for_test(&state, "Background Snapshot").await;
|
|
let (task_id, owner) =
|
|
create_background_task_for_test(&state, &project_id, "web background").await;
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({
|
|
"command": "list_background_tasks",
|
|
"args": { "projectId": project_id, "agentId": owner.to_string() }
|
|
}),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::OK);
|
|
let tasks = body.as_array().expect("background task list is an array");
|
|
assert_eq!(tasks.len(), 1);
|
|
assert_eq!(tasks[0]["taskId"], task_id.to_string());
|
|
assert_eq!(tasks[0]["ownerAgentId"], owner.to_string());
|
|
assert_eq!(tasks[0]["projectId"], project_id);
|
|
assert_eq!(tasks[0]["kind"], "command");
|
|
assert_eq!(tasks[0]["state"], "queued");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn background_actions_are_allowlisted_with_auth_gate() {
|
|
let state = state();
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
let bad_task_id = "not-a-task-id";
|
|
|
|
for command in ["cancel_background_task", "retry_background_task"] {
|
|
let response = request(
|
|
Arc::clone(&state),
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": command, "args": { "taskId": bad_task_id } }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::BAD_REQUEST);
|
|
assert_eq!(
|
|
body["code"], "INVALID",
|
|
"{command} must be an explicit action, not UNKNOWN_COMMAND"
|
|
);
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn mutation_and_pty_commands_stay_out_of_readonly_allowlist() {
|
|
let state = state();
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
for command in ["create_project", "open_terminal", "launch_agent"] {
|
|
let response = request(
|
|
Arc::clone(&state),
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": command, "args": {} }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::BAD_REQUEST);
|
|
assert_eq!(
|
|
body["code"], "UNKNOWN_COMMAND",
|
|
"{command} must stay blocked"
|
|
);
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn invalid_project_id_is_mapped_to_error_dto() {
|
|
let state = state();
|
|
let cookie = pair_and_cookie(Arc::clone(&state)).await;
|
|
|
|
let response = request(
|
|
state,
|
|
Method::POST,
|
|
"/api/invoke",
|
|
json!({ "command": "open_project", "args": { "projectId": "nope" } }),
|
|
&[("cookie", &cookie)],
|
|
)
|
|
.await;
|
|
let (status, body, _) = response_json(response).await;
|
|
|
|
assert_eq!(status, StatusCode::BAD_REQUEST);
|
|
assert_eq!(body["code"], "INVALID");
|
|
}
|
|
|
|
#[test]
|
|
fn direct_project_id_parser_still_rejects_invalid_ids() {
|
|
assert!(parse_project_id("not-a-uuid").is_err());
|
|
}
|
|
}
|