#[cfg(unix)]
mod sessions;
#[cfg(unix)]
mod startup;
#[cfg(unix)]
mod transport;
#[cfg(unix)]
mod turns;
#[cfg(unix)]
use serde::Deserialize;
#[cfg(unix)]
use serde_json::{Value, json};
#[cfg(unix)]
use std::collections::HashSet;
#[cfg(unix)]
const PROTOCOL: &str = "magi-code.local-agent";
#[cfg(unix)]
const VERSION: u32 = 1;
#[cfg(unix)]
const MAX_FRAME_BYTES: usize = 65_536;
#[cfg(unix)]
const MAX_REQUESTS: usize = 4096;
#[cfg(unix)]
const CAPABILITIES: [&str; 9] = [
"initialize",
"status",
"shutdown",
"session.create",
"session.open",
"history.page",
"turn.start",
"turn.cancel",
"approval.answer",
];
#[derive(Debug)]
pub(crate) struct TransportStopped;
impl std::fmt::Display for TransportStopped {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("local-agent transport stopped")
}
}
impl std::error::Error for TransportStopped {}
pub(crate) fn run() -> anyhow::Result<()> {
#[cfg(unix)]
{
transport::run().map_err(|_| TransportStopped.into())
}
#[cfg(not(unix))]
anyhow::bail!("serve --stdio currently requires Unix pipes")
}
#[cfg(unix)]
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct Request {
protocol: String,
version: u32,
id: String,
method: String,
params: Box<serde_json::value::RawValue>,
}
#[cfg(unix)]
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct Initialize {
required_capabilities: Vec<String>,
}
#[cfg(unix)]
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct EmptyParams {}
#[cfg(unix)]
fn empty_params(params: &serde_json::value::RawValue) -> bool {
params.get().starts_with('{') && serde_json::from_str::<EmptyParams>(params.get()).is_ok()
}
#[cfg(unix)]
struct Connection {
instance_id: String,
cwd: std::path::PathBuf,
context_id: Option<String>,
selection: Option<crate::providers::ProviderSelection>,
startup: Option<startup::StartupWorker>,
runtime: Option<startup::RuntimeContext>,
session_worker: Option<sessions::SessionWorker>,
turn_worker: Option<turns::TurnWorker>,
session_request_id: Option<String>,
session_id: Option<String>,
initialize_request_id: Option<String>,
startup_error: Option<&'static str>,
initialized: bool,
request_ids: HashSet<String>,
sequence: u64,
}
#[cfg(unix)]
impl Connection {
fn new() -> anyhow::Result<Self> {
let cwd = std::env::current_dir()?.canonicalize()?;
Ok(Self {
instance_id: uuid::Uuid::new_v4().to_string(),
cwd,
context_id: None,
selection: None,
startup: None,
runtime: None,
session_worker: None,
turn_worker: None,
session_request_id: None,
session_id: None,
initialize_request_id: None,
startup_error: None,
initialized: false,
request_ids: HashSet::new(),
sequence: 0,
})
}
fn response(&mut self, id: Option<&str>, result: Result<Value, &'static str>) -> Vec<u8> {
self.sequence += 1;
let mut frame = json!({
"protocol": PROTOCOL, "version": VERSION, "type": "response",
"instance_id": self.instance_id, "seq": self.sequence, "id": id,
});
match result {
Ok(value) => frame["result"] = value,
Err(code) => frame["error"] = json!({"code": code}),
}
let mut bytes = serde_json::to_vec(&frame).expect("wire values are serializable");
bytes.push(b'\n');
bytes
}
fn receive(&mut self, frame: &[u8]) -> (Option<Vec<u8>>, bool) {
if !bounded_json_depth(frame) {
return (Some(self.response(None, Err("invalid_frame"))), false);
}
if frame.iter().find(|byte| !byte.is_ascii_whitespace()) != Some(&b'{') {
return (Some(self.response(None, Err("invalid_request"))), false);
}
let request: Request = match serde_json::from_slice(frame) {
Ok(request) => request,
Err(_) => return (Some(self.response(None, Err("invalid_request"))), false),
};
if !valid_id(&request.id) || !valid_id(&request.method) {
return (Some(self.response(None, Err("invalid_request"))), false);
}
let id = Some(request.id.as_str());
if request.protocol != PROTOCOL || request.version != VERSION {
return (Some(self.response(id, Err("protocol_mismatch"))), true);
}
if self.request_ids.contains(&request.id) {
return (Some(self.response(id, Err("duplicate_id"))), false);
}
if self.request_ids.len() == MAX_REQUESTS {
return (Some(self.response(id, Err("request_limit"))), true);
}
self.request_ids.insert(request.id.clone());
let result = match request.method.as_str() {
"shutdown" if empty_params(&request.params) => {
return (
Some(self.response(id, Ok(json!({"state": "stopping"})))),
true,
);
}
"initialize" => match self.initialize(&request.params) {
Ok(()) => {
self.initialize_request_id = Some(request.id.clone());
return (None, false);
}
Err(code) => Err(code),
},
"session.create" | "session.open" | "history.page" => {
match self.start_session_request(&request.method, &request.params) {
Ok(()) => {
self.session_request_id = Some(request.id.clone());
return (None, false);
}
Err(code) => Err(code),
}
}
"turn.start" => self.start_turn(&request.params),
"turn.cancel" => self.cancel_turn(&request.params),
"approval.answer" => self.answer_approval(&request.params),
"status" if !self.initialized => Err("not_initialized"),
"status" if empty_params(&request.params) => Ok(self.status()),
"status" | "shutdown" => Err("invalid_params"),
_ => Err("unsupported_method"),
};
(Some(self.response(id, result)), false)
}
fn initialize(&mut self, params: &serde_json::value::RawValue) -> Result<(), &'static str> {
if self.initialized {
return Err("already_initialized");
}
if !params.get().starts_with('{') {
return Err("invalid_params");
}
let params: Initialize =
serde_json::from_str(params.get()).map_err(|_| "invalid_params")?;
if params.required_capabilities.len() > 16
|| params
.required_capabilities
.iter()
.any(|name| !valid_id(name))
{
return Err("invalid_params");
}
if params
.required_capabilities
.iter()
.any(|name| !CAPABILITIES.contains(&name.as_str()))
{
return Err("unsupported_capability");
}
self.startup =
Some(startup::StartupWorker::start(self.cwd.clone()).map_err(|_| "startup_failed")?);
self.initialized = true;
Ok(())
}
fn initialization_result(&self) -> Value {
json!({
"capabilities": CAPABILITIES,
"limits": {
"frame_bytes": MAX_FRAME_BYTES, "json_depth": 32, "id_bytes": 64,
"prompt_bytes": 16_384, "history_page_bytes": 32_768,
"history_page_entries": 100, "content_chunk_bytes": 16_384,
"pending_responses": 32, "requests_per_child": MAX_REQUESTS,
"output_deadline_ms": 2000,
},
"status": self.status(),
})
}
fn status(&self) -> Value {
json!({
"state": if self.initialize_request_id.is_some() { "starting" } else if self.startup_error.is_some() { "failed" } else if self.turn_worker.is_some() { "running" } else { "idle" },
"readiness": if self.initialize_request_id.is_some() { "initializing" } else { self.startup_error.unwrap_or("ready") },
"provider": self.selection.as_ref().and_then(|selection| safe_name(&selection.provider)),
"model": self.selection.as_ref().and_then(|selection| safe_name(&selection.model)),
"context_id": self.context_id, "session_id": self.session_id, "run_id": self.turn_worker.as_ref().map(|w| &w.run_id),
})
}
fn poll_startup(&mut self) -> Option<Vec<u8>> {
let loaded = self.startup.as_mut()?.poll()?;
self.context_id = loaded.context_id;
self.selection = loaded.selection;
let id = self.initialize_request_id.take()?;
let result = match loaded.result {
Ok(runtime) => {
self.runtime = Some(runtime);
Ok(self.initialization_result())
}
Err(code) => {
self.startup_error = Some(code);
Err(code)
}
};
Some(self.response(Some(&id), result))
}
fn start_session_request(
&mut self,
method: &str,
params: &serde_json::value::RawValue,
) -> Result<(), &'static str> {
if !self.initialized || self.initialize_request_id.is_some() {
return Err("not_initialized");
}
if let Some(code) = self.startup_error {
return Err(code);
}
if self.session_worker.is_some() {
return Err("busy");
}
let action = sessions::Action::parse(method, params)?;
let runtime = self.runtime.take().ok_or("busy")?;
match sessions::SessionWorker::start(runtime, action) {
Ok(worker) => {
self.session_worker = Some(worker);
Ok(())
}
Err((runtime, code)) => {
self.runtime = Some(*runtime);
Err(code)
}
}
}
fn poll_session(&mut self) -> Option<Vec<u8>> {
let (runtime, result) = self.session_worker.as_mut()?.poll()?;
self.session_worker = None;
self.session_id = runtime
.attachment
.as_ref()
.map(|a| a.session.id().to_owned());
self.selection = Some(
crate::providers::ProviderSelection::from_config_without_auth(
runtime.effective_config(),
),
);
self.runtime = Some(runtime);
let id = self.session_request_id.take()?;
Some(self.response(Some(&id), result))
}
fn require_ready(&self) -> Result<(), &'static str> {
if !self.initialized || self.initialize_request_id.is_some() {
return Err("not_initialized");
}
if let Some(code) = self.startup_error {
return Err(code);
}
Ok(())
}
fn start_turn(&mut self, params: &serde_json::value::RawValue) -> Result<Value, &'static str> {
self.require_ready()?;
let params: turns::StartParams = turns::parse(params)?;
if crate::sessions::validate_session_id(params.session_id.clone()).is_err()
|| params.session_id.len() > 64
|| params.prompt.len() > 16_384
|| params.prompt.trim().is_empty()
{
return Err("invalid_params");
}
if self.turn_worker.is_some() || self.session_worker.is_some() {
return Err("busy");
}
let attached = self.session_id.as_ref().ok_or("session_missing")?;
if attached != ¶ms.session_id {
return Err("stale_session");
}
let runtime = self.runtime.take().ok_or("busy")?;
match turns::TurnWorker::start(runtime, params.prompt) {
Ok(worker) => {
let result = json!({"run_id":worker.run_id,"accepted":true});
self.turn_worker = Some(worker);
Ok(result)
}
Err((runtime, code)) => {
self.runtime = Some(*runtime);
Err(code)
}
}
}
fn active_run(
&self,
session_id: &str,
run_id: &str,
) -> Result<&turns::TurnWorker, &'static str> {
self.require_ready()?;
if !valid_id(session_id) || !valid_id(run_id) {
return Err("invalid_params");
}
let worker = self.turn_worker.as_ref().ok_or("stale_run")?;
if worker.session_id != session_id || worker.run_id != run_id {
return Err("stale_run");
}
Ok(worker)
}
fn cancel_turn(&self, params: &serde_json::value::RawValue) -> Result<Value, &'static str> {
let params: turns::RunParams = turns::parse(params)?;
self.active_run(¶ms.session_id, ¶ms.run_id)?
.cancel();
Ok(json!({"cancellation_requested":true}))
}
fn answer_approval(&self, params: &serde_json::value::RawValue) -> Result<Value, &'static str> {
let params: turns::AnswerParams = turns::parse(params)?;
if !valid_id(¶ms.approval_id) {
return Err("invalid_params");
}
self.active_run(¶ms.session_id, ¶ms.run_id)
.map_err(|code| {
if code == "stale_run" {
"stale_approval"
} else {
code
}
})?
.answer(¶ms.approval_id, params.allow)?;
Ok(json!({"recorded":true}))
}
fn event(&mut self, name: &str, data: Value) -> Vec<u8> {
self.sequence += 1;
let mut bytes = serde_json::to_vec(&json!({"protocol":PROTOCOL,"version":VERSION,"type":"event","instance_id":self.instance_id,"seq":self.sequence,"event":name,"data":data})).expect("wire values serializable");
bytes.push(b'\n');
bytes
}
fn poll_turn(&mut self) -> Option<Vec<u8>> {
let worker = self.turn_worker.as_mut()?;
if let Some(event) = worker.next_event() {
return Some(self.event(event.name, event.data));
}
let (runtime, mut terminal) = worker.poll()?;
terminal["session_id"] = json!(worker.session_id);
terminal["run_id"] = json!(worker.run_id);
self.turn_worker = None;
self.runtime = Some(runtime);
Some(self.event("run.finished", terminal))
}
fn cancel_startup(&self) {
if let Some(worker) = &self.startup {
worker.cancel();
}
if let Some(worker) = &self.session_worker {
worker.cancel();
}
if let Some(worker) = &self.turn_worker {
worker.cancel();
}
}
fn begin_shutdown(&mut self) -> Option<Vec<u8>> {
self.cancel_startup();
if let Some(id) = self.initialize_request_id.take() {
Some(self.response(Some(&id), Err("startup_cancelled")))
} else {
let id = self.session_request_id.take()?;
Some(self.response(Some(&id), Err("operation_cancelled")))
}
}
}
#[cfg(unix)]
impl Drop for Connection {
fn drop(&mut self) {
if let Some(worker) = self.startup.take() {
worker.cleanup(
self.runtime.take(),
self.session_worker.take(),
self.turn_worker.take(),
std::time::Instant::now() + std::time::Duration::from_secs(2),
);
}
}
}
#[cfg(unix)]
fn safe_name(name: &str) -> Option<String> {
let redacted = crate::output::redact_sensitive_text(name);
(name.len() <= 256
&& !name.chars().any(char::is_control)
&& redacted == name
&& !crate::config::looks_like_secret_value(name))
.then(|| name.to_owned())
}
#[cfg(unix)]
fn valid_id(value: &str) -> bool {
!value.is_empty()
&& value.len() <= 64
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || b"._-".contains(&byte))
}
#[cfg(unix)]
fn bounded_json_depth(bytes: &[u8]) -> bool {
let (mut depth, mut quoted, mut escaped) = (0usize, false, false);
for &byte in bytes {
if quoted {
if escaped {
escaped = false;
} else if byte == b'\\' {
escaped = true;
} else if byte == b'"' {
quoted = false;
}
} else {
match byte {
b'"' => quoted = true,
b'{' | b'[' => {
depth += 1;
if depth > 32 {
return false;
}
}
b'}' | b']' => {
let Some(next) = depth.checked_sub(1) else {
return false;
};
depth = next;
}
_ => {}
}
}
}
depth == 0 && !quoted
}