use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use futures::{SinkExt, StreamExt};
use supercode_harness::{
CoordinatedRuntime, RuntimeAuthorization, RuntimeClientId, SdkRuntime,
DEFAULT_RUNTIME_LEASE_TTL_MS,
};
use tokio::net::{TcpListener, TcpStream};
#[cfg(unix)]
use tokio::net::{UnixListener, UnixStream};
use tokio::sync::{watch, Notify};
use tokio_tungstenite::tungstenite::handshake::server::{ErrorResponse, Request, Response};
use tokio_tungstenite::tungstenite::http::StatusCode;
use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
use tokio_tungstenite::tungstenite::Message;
use zeroize::Zeroize;
use zeroize::Zeroizing;
use crate::codex_app_server_v0_144::{CodexAppServerAdapter, CodexCompatibilityMode};
const TOKEN_BYTES: usize = 32;
const MAX_MESSAGE_BYTES: usize = 16 * 1024 * 1024;
const MAX_FRAME_BYTES: usize = 1024 * 1024;
const MAX_WRITE_BUFFER_BYTES: usize = 2 * MAX_MESSAGE_BYTES;
const MAX_HANDSHAKE_BYTES: usize = 64 * 1024;
const MAX_CONCURRENT_CONNECTIONS: usize = 32;
const HANDSHAKE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
#[derive(Debug, thiserror::Error)]
pub enum CredentialError {
#[error("operating system random source failed")]
RandomSource,
#[error("credential generation collided repeatedly")]
Collision,
#[error("credential is already revoked")]
Revoked,
#[error("bootstrap credential is invalid for this runtime generation")]
InvalidBootstrap,
#[error("invalid runtime client identity: {0}")]
InvalidClient(String),
#[error("private credential file operation failed: {0}")]
Io(#[from] std::io::Error),
#[error("endpoint configuration is already shared")]
ConfigurationLocked,
}
struct CredentialGrant {
client_id: RuntimeClientId,
authorization: RuntimeAuthorization,
revoked: watch::Receiver<bool>,
runtime_id: String,
generation: [u8; 16],
active_channel: watch::Sender<bool>,
}
impl Drop for CredentialGrant {
fn drop(&mut self) {
self.active_channel.send_replace(false);
}
}
struct CredentialRecord {
client_id: RuntimeClientId,
authorization: RuntimeAuthorization,
revoked: watch::Sender<bool>,
active_channel: watch::Sender<bool>,
}
struct CodexCredentialRegistry {
records: Mutex<BTreeMap<[u8; 32], CredentialRecord>>,
runtime_id: String,
generation: [u8; 16],
}
pub struct CodexBootstrapCredential {
secret: [u8; TOKEN_BYTES],
digest: [u8; 32],
generation: [u8; 16],
revoked: bool,
}
impl fmt::Debug for CodexBootstrapCredential {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("CodexBootstrapCredential([REDACTED])")
}
}
impl Drop for CodexBootstrapCredential {
fn drop(&mut self) {
zero(&mut self.secret);
}
}
pub struct CodexBootstrapFile {
path: PathBuf,
generation: [u8; 16],
}
impl Drop for CodexBootstrapFile {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.path);
}
}
impl CodexCredentialRegistry {
fn new(runtime_id: String, generation: [u8; 16]) -> Arc<Self> {
Arc::new(Self {
records: Mutex::new(BTreeMap::new()),
runtime_id,
generation,
})
}
fn issue(
self: &Arc<Self>,
client_id: impl Into<String>,
authorization: RuntimeAuthorization,
) -> Result<CodexClientCredential, CredentialError> {
self.issue_with_fill(client_id, authorization, |secret| {
getrandom::getrandom(secret).map_err(|_| CredentialError::RandomSource)
})
}
fn issue_with_fill(
self: &Arc<Self>,
client_id: impl Into<String>,
authorization: RuntimeAuthorization,
mut fill: impl FnMut(&mut [u8; TOKEN_BYTES]) -> Result<(), CredentialError>,
) -> Result<CodexClientCredential, CredentialError> {
let client_id = RuntimeClientId::parse(client_id.into())
.map_err(|error| CredentialError::InvalidClient(error.to_string()))?;
for _ in 0..3 {
let mut secret = [0u8; TOKEN_BYTES];
fill(&mut secret)?;
let digest = credential_digest(&secret);
let mut records = self
.records
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if records.contains_key(&digest) {
zero(&mut secret);
continue;
}
let (revoked, _) = watch::channel(false);
let (active_channel, _) = watch::channel(false);
records.insert(
digest,
CredentialRecord {
client_id: client_id.clone(),
authorization,
revoked,
active_channel,
},
);
return Ok(CodexClientCredential {
secret,
digest,
client_id,
revoked: false,
});
}
Err(CredentialError::Collision)
}
async fn revoke(&self, credential: &mut CodexClientCredential) -> Result<(), CredentialError> {
if credential.revoked {
return Err(CredentialError::Revoked);
}
let record = self
.records
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(&credential.digest)
.ok_or(CredentialError::Revoked)?;
let mut active_channel = record.active_channel.subscribe();
let _ = record.revoked.send(true);
while *active_channel.borrow() {
if active_channel.changed().await.is_err() {
break;
}
}
credential.revoked = true;
zero(&mut credential.secret);
Ok(())
}
async fn rotate(
self: &Arc<Self>,
credential: &mut CodexClientCredential,
client_id: impl Into<String>,
authorization: RuntimeAuthorization,
) -> Result<CodexClientCredential, CredentialError> {
let mut replacement = self.issue(client_id, authorization)?;
if let Err(error) = self.revoke(credential).await {
let _ = self.revoke(&mut replacement).await;
return Err(error);
}
Ok(replacement)
}
async fn revoke_all(&self) {
let records = {
let mut records = self
.records
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
std::mem::take(&mut *records)
};
let mut channels = Vec::with_capacity(records.len());
for (_, record) in records {
let mut active = record.active_channel.subscribe();
let _ = record.revoked.send(true);
channels.push(async move {
while *active.borrow() {
if active.changed().await.is_err() {
break;
}
}
});
}
futures::future::join_all(channels).await;
}
fn authenticate(&self, token: &str) -> Result<CredentialGrant, AuthenticationError> {
let (presented_client_id, secret_text) = token
.split_once('.')
.ok_or(AuthenticationError::Unauthenticated)?;
let mut secret =
decode_hex_secret(secret_text).ok_or(AuthenticationError::Unauthenticated)?;
let digest = credential_digest(&secret);
zero(&mut secret);
let mut records = self
.records
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let record = records
.get_mut(&digest)
.ok_or(AuthenticationError::Unauthenticated)?;
if record.client_id.as_str() != presented_client_id {
return Err(AuthenticationError::Unauthorized);
}
if *record.active_channel.borrow() {
return Err(AuthenticationError::Busy);
}
record.active_channel.send_replace(true);
Ok(CredentialGrant {
client_id: record.client_id.clone(),
authorization: record.authorization.clone(),
revoked: record.revoked.subscribe(),
runtime_id: self.runtime_id.clone(),
generation: self.generation,
active_channel: record.active_channel.clone(),
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum AuthenticationError {
Unauthenticated,
Unauthorized,
Busy,
}
pub struct CodexClientCredential {
secret: [u8; TOKEN_BYTES],
digest: [u8; 32],
client_id: RuntimeClientId,
revoked: bool,
}
impl CodexClientCredential {
pub fn spawn_tokio_child(
&self,
command: &mut tokio::process::Command,
variable: &str,
) -> std::io::Result<tokio::process::Child> {
let bearer = Zeroizing::new(self.bearer_value());
command.env(variable, bearer.as_str());
let child = command.spawn();
command.env_remove(variable);
child
}
fn bearer_value(&self) -> String {
format!("{}.{}", self.client_id.as_str(), encode_hex(&self.secret))
}
}
impl fmt::Debug for CodexClientCredential {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("CodexClientCredential([REDACTED])")
}
}
impl Drop for CodexClientCredential {
fn drop(&mut self) {
zero(&mut self.secret);
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CodexEndpointKind {
WebSocket(SocketAddr),
#[cfg(unix)]
Unix(PathBuf),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CodexEndpointHealth {
Ready,
ShuttingDown,
Stopped,
}
pub struct CodexEndpoint {
coordinator: Arc<CoordinatedRuntime>,
credentials: Arc<CodexCredentialRegistry>,
bootstrap_digest: Mutex<[u8; 32]>,
cwd: PathBuf,
client_home: Option<PathBuf>,
mode: CodexCompatibilityMode,
forwarded_ports: BTreeSet<u16>,
}
impl CodexEndpoint {
pub fn new(
runtime: Arc<dyn SdkRuntime>,
runtime_id: impl Into<String>,
cwd: impl Into<PathBuf>,
) -> Result<(Arc<Self>, CodexBootstrapCredential), CredentialError> {
let generation = random_generation()?;
let bootstrap = new_bootstrap(generation)?;
let endpoint = Arc::new(Self {
coordinator: CoordinatedRuntime::with_lease_ttl(runtime, DEFAULT_RUNTIME_LEASE_TTL_MS),
credentials: CodexCredentialRegistry::new(runtime_id.into(), generation),
bootstrap_digest: Mutex::new(bootstrap.digest),
cwd: cwd.into(),
client_home: None,
mode: CodexCompatibilityMode::TracedOnly,
forwarded_ports: BTreeSet::new(),
});
Ok((endpoint, bootstrap))
}
pub fn with_mode(
runtime: Arc<dyn SdkRuntime>,
runtime_id: impl Into<String>,
cwd: impl Into<PathBuf>,
mode: CodexCompatibilityMode,
) -> Result<(Arc<Self>, CodexBootstrapCredential), CredentialError> {
let generation = random_generation()?;
let bootstrap = new_bootstrap(generation)?;
let endpoint = Arc::new(Self {
coordinator: CoordinatedRuntime::with_lease_ttl(runtime, DEFAULT_RUNTIME_LEASE_TTL_MS),
credentials: CodexCredentialRegistry::new(runtime_id.into(), generation),
bootstrap_digest: Mutex::new(bootstrap.digest),
cwd: cwd.into(),
client_home: None,
mode,
forwarded_ports: BTreeSet::new(),
});
Ok((endpoint, bootstrap))
}
pub fn with_forwarded_ports(
mut self: Arc<Self>,
ports: impl IntoIterator<Item = u16>,
) -> Result<Arc<Self>, CredentialError> {
Arc::get_mut(&mut self)
.ok_or(CredentialError::ConfigurationLocked)?
.forwarded_ports
.extend(ports.into_iter().filter(|port| *port != 0));
Ok(self)
}
pub fn with_client_home(
mut self: Arc<Self>,
path: impl Into<PathBuf>,
) -> Result<Arc<Self>, CredentialError> {
Arc::get_mut(&mut self)
.ok_or(CredentialError::ConfigurationLocked)?
.client_home = Some(path.into());
Ok(self)
}
pub fn issue_interactive(
self: &Arc<Self>,
bootstrap: &CodexBootstrapCredential,
client_id: impl Into<String>,
) -> Result<CodexClientCredential, CredentialError> {
self.authenticate_bootstrap(bootstrap)?;
self.credentials.issue(
client_id,
RuntimeAuthorization::new([
supercode_harness::RuntimePermission::Observe,
supercode_harness::RuntimePermission::Interact,
supercode_harness::RuntimePermission::Approve,
]),
)
}
pub fn issue_observer(
self: &Arc<Self>,
bootstrap: &CodexBootstrapCredential,
client_id: impl Into<String>,
) -> Result<CodexClientCredential, CredentialError> {
self.authenticate_bootstrap(bootstrap)?;
self.credentials
.issue(client_id, RuntimeAuthorization::observer())
}
pub async fn rotate_interactive(
self: &Arc<Self>,
bootstrap: &CodexBootstrapCredential,
credential: &mut CodexClientCredential,
client_id: impl Into<String>,
) -> Result<CodexClientCredential, CredentialError> {
self.authenticate_bootstrap(bootstrap)?;
self.credentials
.rotate(
credential,
client_id,
RuntimeAuthorization::new([
supercode_harness::RuntimePermission::Observe,
supercode_harness::RuntimePermission::Interact,
supercode_harness::RuntimePermission::Approve,
]),
)
.await
}
pub fn rotate_bootstrap(
&self,
bootstrap: &mut CodexBootstrapCredential,
) -> Result<CodexBootstrapCredential, CredentialError> {
let replacement = new_bootstrap(self.credentials.generation)?;
let mut expected = self
.bootstrap_digest
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if bootstrap.revoked
|| bootstrap.generation != self.credentials.generation
|| *expected != bootstrap.digest
|| credential_digest(&bootstrap.secret) != *expected
{
return Err(CredentialError::InvalidBootstrap);
}
*expected = replacement.digest;
bootstrap.revoked = true;
zero(&mut bootstrap.secret);
Ok(replacement)
}
#[cfg(unix)]
pub fn rotate_bootstrap_file(
&self,
bootstrap: &mut CodexBootstrapCredential,
bootstrap_file: &mut CodexBootstrapFile,
runtime_dir: impl Into<PathBuf>,
) -> Result<(CodexBootstrapCredential, CodexBootstrapFile), CredentialError> {
self.rotate_bootstrap_file_with_remove(
bootstrap,
bootstrap_file,
runtime_dir.into(),
|path| std::fs::remove_file(path),
)
}
#[cfg(unix)]
fn rotate_bootstrap_file_with_remove(
&self,
bootstrap: &mut CodexBootstrapCredential,
bootstrap_file: &mut CodexBootstrapFile,
runtime_dir: PathBuf,
remove_old: impl FnOnce(&std::path::Path) -> std::io::Result<()>,
) -> Result<(CodexBootstrapCredential, CodexBootstrapFile), CredentialError> {
self.authenticate_bootstrap(bootstrap)?;
self.authenticate_bootstrap_file(bootstrap_file)?;
let replacement = new_bootstrap(self.credentials.generation)?;
let replacement_file =
write_bootstrap_secret(runtime_dir, &replacement.secret, replacement.generation)?;
let mut expected = self
.bootstrap_digest
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if *expected != bootstrap.digest || credential_digest(&bootstrap.secret) != *expected {
drop(replacement_file);
return Err(CredentialError::InvalidBootstrap);
}
remove_old(&bootstrap_file.path)?;
*expected = replacement.digest;
bootstrap.revoked = true;
zero(&mut bootstrap.secret);
bootstrap_file.generation = [0; 16];
Ok((replacement, replacement_file))
}
pub async fn revoke_frontend(
&self,
bootstrap: &CodexBootstrapCredential,
credential: &mut CodexClientCredential,
) -> Result<(), CredentialError> {
self.authenticate_bootstrap(bootstrap)?;
self.credentials.revoke(credential).await
}
#[cfg(unix)]
pub fn write_bootstrap_file(
&self,
bootstrap: &CodexBootstrapCredential,
runtime_dir: impl Into<PathBuf>,
) -> Result<CodexBootstrapFile, CredentialError> {
self.authenticate_bootstrap(bootstrap)?;
write_bootstrap_secret(runtime_dir.into(), &bootstrap.secret, bootstrap.generation)
}
#[cfg(unix)]
pub fn issue_interactive_from_file(
self: &Arc<Self>,
bootstrap_file: &CodexBootstrapFile,
client_id: impl Into<String>,
) -> Result<CodexClientCredential, CredentialError> {
self.authenticate_bootstrap_file(bootstrap_file)?;
self.credentials.issue(
client_id,
RuntimeAuthorization::new([
supercode_harness::RuntimePermission::Observe,
supercode_harness::RuntimePermission::Interact,
supercode_harness::RuntimePermission::Approve,
]),
)
}
#[cfg(unix)]
pub async fn revoke_frontend_from_file(
&self,
bootstrap_file: &CodexBootstrapFile,
credential: &mut CodexClientCredential,
) -> Result<(), CredentialError> {
self.authenticate_bootstrap_file(bootstrap_file)?;
self.credentials.revoke(credential).await
}
pub async fn terminate(
self: &Arc<Self>,
bootstrap: &CodexBootstrapCredential,
) -> Result<(), CredentialError> {
self.authenticate_bootstrap(bootstrap)?;
self.credentials.revoke_all().await;
let client_id = RuntimeClientId::parse("supercode-bootstrap-operator")
.map_err(|error| CredentialError::InvalidClient(error.to_string()))?;
self.coordinator
.client(client_id, RuntimeAuthorization::owner())
.close()
.await
.map_err(|_| CredentialError::InvalidBootstrap)
}
fn authenticate_bootstrap(
&self,
bootstrap: &CodexBootstrapCredential,
) -> Result<(), CredentialError> {
if bootstrap.revoked || bootstrap.generation != self.credentials.generation {
return Err(CredentialError::InvalidBootstrap);
}
let expected = self
.bootstrap_digest
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if *expected != bootstrap.digest || credential_digest(&bootstrap.secret) != *expected {
return Err(CredentialError::InvalidBootstrap);
}
Ok(())
}
#[cfg(unix)]
fn authenticate_bootstrap_file(
&self,
bootstrap_file: &CodexBootstrapFile,
) -> Result<(), CredentialError> {
if bootstrap_file.generation != self.credentials.generation {
return Err(CredentialError::InvalidBootstrap);
}
let mut encoded = std::fs::read(&bootstrap_file.path)?;
let decoded = std::str::from_utf8(&encoded)
.ok()
.and_then(decode_hex_secret)
.ok_or(CredentialError::InvalidBootstrap);
encoded.zeroize();
let mut secret = decoded?;
let digest = credential_digest(&secret);
zero(&mut secret);
let expected = self
.bootstrap_digest
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if digest != *expected {
return Err(CredentialError::InvalidBootstrap);
}
Ok(())
}
pub async fn bind_websocket(self: &Arc<Self>, port: u16) -> std::io::Result<CodexServerHandle> {
let listener =
TcpListener::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), port)).await?;
let address = listener.local_addr()?;
let shutdown = Arc::new(Notify::new());
let shutting_down = Arc::new(AtomicBool::new(false));
let stopped = Arc::new(AtomicBool::new(false));
let task = tokio::spawn(run_tcp_accept_loop(
listener,
self.clone(),
address.to_string(),
self.forwarded_ports.clone(),
shutdown.clone(),
shutting_down.clone(),
stopped.clone(),
));
Ok(CodexServerHandle {
kind: CodexEndpointKind::WebSocket(address),
shutdown,
shutting_down,
stopped,
task: Some(task),
})
}
#[cfg(unix)]
pub async fn bind_unix(
self: &Arc<Self>,
socket_path: impl Into<PathBuf>,
) -> std::io::Result<CodexServerHandle> {
use std::os::unix::fs::{DirBuilderExt, MetadataExt, PermissionsExt};
let socket_path = socket_path.into();
let parent = socket_path.parent().ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::InvalidInput, "socket has no parent")
})?;
match std::fs::symlink_metadata(parent) {
Ok(metadata) => {
if metadata.file_type().is_symlink() || !metadata.is_dir() {
return Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"Unix endpoint directory must be a real directory",
));
}
if metadata.permissions().mode() & 0o777 != 0o700 {
return Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"Unix endpoint directory must have mode 0700",
));
}
if metadata.uid() != unsafe { libc::geteuid() } {
return Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"Unix endpoint directory must be owned by this user",
));
}
}
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
let mut builder = std::fs::DirBuilder::new();
builder.mode(0o700).create(parent)?;
}
Err(error) => return Err(error),
}
match std::fs::symlink_metadata(&socket_path) {
Ok(_) => {
return Err(std::io::Error::new(
std::io::ErrorKind::AlreadyExists,
"refusing to replace an existing Unix endpoint",
))
}
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {}
Err(error) => return Err(error),
}
let listener = UnixListener::bind(&socket_path)?;
std::fs::set_permissions(&socket_path, std::fs::Permissions::from_mode(0o600))?;
let shutdown = Arc::new(Notify::new());
let shutting_down = Arc::new(AtomicBool::new(false));
let stopped = Arc::new(AtomicBool::new(false));
let task = tokio::spawn(run_unix_accept_loop(
listener,
self.clone(),
socket_path.clone(),
shutdown.clone(),
shutting_down.clone(),
stopped.clone(),
));
Ok(CodexServerHandle {
kind: CodexEndpointKind::Unix(socket_path),
shutdown,
shutting_down,
stopped,
task: Some(task),
})
}
fn authenticated_adapter(&self, grant: &CredentialGrant) -> Arc<CodexAppServerAdapter> {
debug_assert_eq!(grant.runtime_id, self.credentials.runtime_id);
debug_assert_eq!(grant.generation, self.credentials.generation);
let runtime = self
.coordinator
.client(grant.client_id.clone(), grant.authorization.clone());
CodexAppServerAdapter::with_mode_and_client_home(
runtime,
self.cwd.clone(),
self.client_home.clone().unwrap_or_else(|| self.cwd.clone()),
self.mode,
)
}
}
pub struct CodexServerHandle {
kind: CodexEndpointKind,
shutdown: Arc<Notify>,
shutting_down: Arc<AtomicBool>,
stopped: Arc<AtomicBool>,
task: Option<tokio::task::JoinHandle<()>>,
}
impl CodexServerHandle {
pub fn kind(&self) -> &CodexEndpointKind {
&self.kind
}
pub fn health(&self) -> CodexEndpointHealth {
if self.stopped.load(Ordering::SeqCst) {
CodexEndpointHealth::Stopped
} else if self.shutting_down.load(Ordering::SeqCst) {
CodexEndpointHealth::ShuttingDown
} else {
CodexEndpointHealth::Ready
}
}
pub async fn shutdown(mut self) {
self.shutting_down.store(true, Ordering::SeqCst);
self.shutdown.notify_waiters();
if let Some(task) = self.task.take() {
let _ = task.await;
}
}
}
impl Drop for CodexServerHandle {
fn drop(&mut self) {
self.shutting_down.store(true, Ordering::SeqCst);
self.shutdown.notify_waiters();
if let Some(task) = self.task.take() {
task.abort();
}
}
}
async fn run_tcp_accept_loop(
listener: TcpListener,
endpoint: Arc<CodexEndpoint>,
expected_host: String,
forwarded_ports: BTreeSet<u16>,
shutdown: Arc<Notify>,
shutting_down: Arc<AtomicBool>,
stopped: Arc<AtomicBool>,
) {
let connection_budget = Arc::new(tokio::sync::Semaphore::new(MAX_CONCURRENT_CONNECTIONS));
loop {
tokio::select! {
_ = shutdown.notified() => break,
accepted = listener.accept() => match accepted {
Ok((stream, _)) => {
let Ok(permit) = connection_budget.clone().try_acquire_owned() else {
drop(stream);
continue;
};
let endpoint = endpoint.clone();
let expected_host = expected_host.clone();
let forwarded_ports = forwarded_ports.clone();
tokio::spawn(async move {
let _permit = permit;
let _ = serve_stream(stream, endpoint, HostRule::Tcp { expected_host, forwarded_ports }).await;
});
}
Err(_) => break,
}
}
}
shutting_down.store(true, Ordering::SeqCst);
stopped.store(true, Ordering::SeqCst);
}
#[cfg(unix)]
async fn run_unix_accept_loop(
listener: UnixListener,
endpoint: Arc<CodexEndpoint>,
socket_path: PathBuf,
shutdown: Arc<Notify>,
shutting_down: Arc<AtomicBool>,
stopped: Arc<AtomicBool>,
) {
let connection_budget = Arc::new(tokio::sync::Semaphore::new(MAX_CONCURRENT_CONNECTIONS));
loop {
tokio::select! {
_ = shutdown.notified() => break,
accepted = listener.accept() => match accepted {
Ok((stream, _)) => {
let Ok(permit) = connection_budget.clone().try_acquire_owned() else {
drop(stream);
continue;
};
let endpoint = endpoint.clone();
tokio::spawn(async move {
let _permit = permit;
let _ = serve_stream(stream, endpoint, HostRule::Unix).await;
});
}
Err(_) => break,
}
}
}
let _ = std::fs::remove_file(socket_path);
shutting_down.store(true, Ordering::SeqCst);
stopped.store(true, Ordering::SeqCst);
}
enum HostRule {
Tcp {
expected_host: String,
forwarded_ports: BTreeSet<u16>,
},
#[cfg(unix)]
Unix,
}
trait LocalStream: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static {}
impl LocalStream for TcpStream {}
#[cfg(unix)]
impl LocalStream for UnixStream {}
struct HandshakeBoundedStream<S> {
inner: S,
bytes_read: usize,
complete: Arc<AtomicBool>,
}
impl<S: Unpin> Unpin for HandshakeBoundedStream<S> {}
impl<S: LocalStream> LocalStream for HandshakeBoundedStream<S> {}
impl<S: tokio::io::AsyncRead + Unpin> tokio::io::AsyncRead for HandshakeBoundedStream<S> {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
buffer: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
let before = buffer.filled().len();
match std::pin::Pin::new(&mut self.inner).poll_read(context, buffer) {
std::task::Poll::Ready(Ok(())) => {
if !self.complete.load(Ordering::Acquire) {
self.bytes_read = self
.bytes_read
.saturating_add(buffer.filled().len().saturating_sub(before));
if self.bytes_read > MAX_HANDSHAKE_BYTES {
return std::task::Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"WebSocket handshake exceeded byte limit",
)));
}
}
std::task::Poll::Ready(Ok(()))
}
other => other,
}
}
}
impl<S: tokio::io::AsyncWrite + Unpin> tokio::io::AsyncWrite for HandshakeBoundedStream<S> {
fn poll_write(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
buffer: &[u8],
) -> std::task::Poll<std::io::Result<usize>> {
std::pin::Pin::new(&mut self.inner).poll_write(context, buffer)
}
fn poll_flush(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::pin::Pin::new(&mut self.inner).poll_flush(context)
}
fn poll_shutdown(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::pin::Pin::new(&mut self.inner).poll_shutdown(context)
}
}
#[allow(clippy::result_large_err)]
async fn serve_stream<S: LocalStream>(
stream: S,
endpoint: Arc<CodexEndpoint>,
host_rule: HostRule,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let selected = Arc::new(Mutex::new(None::<CredentialGrant>));
let callback_selected = selected.clone();
let credentials = endpoint.credentials.clone();
let handshake_complete = Arc::new(AtomicBool::new(false));
let callback_complete = handshake_complete.clone();
let stream = HandshakeBoundedStream {
inner: stream,
bytes_read: 0,
complete: handshake_complete,
};
let websocket = tokio::time::timeout(
HANDSHAKE_TIMEOUT,
tokio_tungstenite::accept_hdr_async_with_config(
stream,
move |request: &Request, response: Response| -> Result<Response, ErrorResponse> {
match validate_upgrade(request, &host_rule, &credentials) {
Ok(grant) => {
*callback_selected
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = Some(grant);
callback_complete.store(true, Ordering::Release);
Ok(response)
}
Err((status, message)) => Err(http_error(status, message)),
}
},
Some(websocket_config()),
),
)
.await
.map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::TimedOut,
"WebSocket handshake timed out",
)
})??;
let grant = selected
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take()
.ok_or("authenticated upgrade did not select a grant")?;
serve_websocket(websocket, endpoint, grant).await
}
async fn serve_websocket<S: LocalStream>(
mut websocket: tokio_tungstenite::WebSocketStream<S>,
endpoint: Arc<CodexEndpoint>,
mut grant: CredentialGrant,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let runtime = endpoint.authenticated_adapter(&grant);
let mut connection = runtime.connection();
loop {
tokio::select! {
changed = grant.revoked.changed() => {
if changed.is_err() || *grant.revoked.borrow() {
let _ = websocket.close(None).await;
break;
}
}
incoming = websocket.next() => {
let Some(message) = incoming else { break; };
match message? {
Message::Text(text) => {
let value: serde_json::Value = match serde_json::from_str(&text) {
Ok(value) => value,
Err(_) => {
websocket.send(Message::Text(serde_json::json!({"id":null,"error":{
"code":-32700,"message":"invalid JSON request","data":{"name":"parse_error"}
}}).to_string().into())).await?;
continue;
}
};
if contains_application_credential(&value) {
let id = value.get("id").cloned().unwrap_or(serde_json::Value::Null);
websocket.send(Message::Text(serde_json::json!({"id":id,"error":{
"code":-32030,"message":"application-message credentials are forbidden",
"data":{"name":"unauthenticated"}
}}).to_string().into())).await?;
continue;
}
let output = if value.get("method").is_some() {
connection.handle(value).await
} else {
connection.handle_server_response(&value).await?
};
for message in output {
websocket.send(Message::Text(message.to_string().into())).await?;
}
}
Message::Close(_) => break,
Message::Ping(payload) => websocket.send(Message::Pong(payload)).await?,
Message::Binary(_) => {
websocket.close(None).await?;
break;
}
Message::Pong(_) | Message::Frame(_) => {}
}
}
notifications = connection.next_notifications(), if connection_is_attached(&connection) => {
for message in notifications? {
websocket.send(Message::Text(message.to_string().into())).await?;
}
}
}
}
let _ = runtime.detach().await;
Ok(())
}
fn connection_is_attached(connection: &crate::CodexConnection) -> bool {
connection.is_attached()
}
fn validate_upgrade(
request: &Request,
host_rule: &HostRule,
credentials: &CodexCredentialRegistry,
) -> Result<CredentialGrant, (StatusCode, &'static str)> {
if request.uri().path() != "/" || request.uri().query().is_some() {
return Err((StatusCode::NOT_FOUND, "unsupported endpoint"));
}
if request.headers().contains_key("cookie")
|| request.headers().contains_key("sec-websocket-protocol")
|| request.headers().contains_key("transfer-encoding")
|| request.headers().contains_key("content-length")
{
return Err((StatusCode::BAD_REQUEST, "forbidden credential carrier"));
}
for name in [
"host",
"origin",
"authorization",
"content-length",
"x-supercode-permissions",
] {
if request.headers().get_all(name).iter().count() > 1 {
return Err((StatusCode::BAD_REQUEST, "duplicate security header"));
}
}
if request.uri().scheme().is_some() || request.uri().authority().is_some() {
return Err((StatusCode::BAD_REQUEST, "absolute-form target rejected"));
}
let host = request
.headers()
.get("host")
.and_then(|value| value.to_str().ok())
.ok_or((StatusCode::BAD_REQUEST, "missing Host"))?;
match host_rule {
HostRule::Tcp {
expected_host,
forwarded_ports,
} => {
let accepted_forward = host
.strip_prefix("127.0.0.1:")
.and_then(|port| port.parse::<u16>().ok())
.is_some_and(|port| forwarded_ports.contains(&port));
if host != expected_host && !accepted_forward {
return Err((StatusCode::BAD_REQUEST, "invalid loopback Host"));
}
}
#[cfg(unix)]
HostRule::Unix if host != "localhost" => {
return Err((StatusCode::BAD_REQUEST, "invalid Unix Host"));
}
#[cfg(unix)]
HostRule::Unix => {}
}
if request.headers().contains_key("origin") {
return Err((StatusCode::BAD_REQUEST, "Origin rejected"));
}
let authorization = request
.headers()
.get("authorization")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "))
.ok_or((StatusCode::UNAUTHORIZED, "authentication failed"))?;
let mut grant = credentials
.authenticate(authorization)
.map_err(|error| match error {
AuthenticationError::Unauthenticated => {
(StatusCode::UNAUTHORIZED, "authentication failed")
}
AuthenticationError::Unauthorized => (StatusCode::FORBIDDEN, "client id rejected"),
AuthenticationError::Busy => (StatusCode::CONFLICT, "client channel busy"),
})?;
if let Some(requested) = request.headers().get("x-supercode-permissions") {
let requested = requested
.to_str()
.ok()
.and_then(|value| RuntimeAuthorization::parse_header(value).ok())
.ok_or((StatusCode::BAD_REQUEST, "invalid permission narrowing"))?;
grant.authorization = grant.authorization.restrict_to(&requested);
}
Ok(grant)
}
fn contains_application_credential(value: &serde_json::Value) -> bool {
match value {
serde_json::Value::Object(object) => object.iter().any(|(key, child)| {
matches!(
key.to_ascii_lowercase().as_str(),
"authorization" | "bearer" | "credential" | "token"
) || contains_application_credential(child)
}),
serde_json::Value::Array(array) => array.iter().any(contains_application_credential),
_ => false,
}
}
fn websocket_config() -> WebSocketConfig {
let mut config = WebSocketConfig::default();
config.write_buffer_size = 128 * 1024;
config.max_write_buffer_size = MAX_WRITE_BUFFER_BYTES;
config.max_message_size = Some(MAX_MESSAGE_BYTES);
config.max_frame_size = Some(MAX_FRAME_BYTES);
config.accept_unmasked_frames = false;
config
}
fn http_error(status: StatusCode, message: &'static str) -> ErrorResponse {
let (code, name) = match status {
StatusCode::UNAUTHORIZED => (-32030, "unauthenticated"),
StatusCode::FORBIDDEN => (-32031, "unauthorized"),
StatusCode::CONFLICT => (-32000, "busy"),
_ => (-32600, "invalid_request"),
};
tokio_tungstenite::tungstenite::http::Response::builder()
.status(status)
.header("content-type", "application/json")
.body(Some(
serde_json::json!({"error":{"code":code,"message":message,"data":{"name":name}}})
.to_string(),
))
.expect("static HTTP error response")
}
fn credential_digest(secret: &[u8; TOKEN_BYTES]) -> [u8; 32] {
*blake3::hash(secret).as_bytes()
}
fn random_generation() -> Result<[u8; 16], CredentialError> {
let mut generation = [0u8; 16];
getrandom::getrandom(&mut generation).map_err(|_| CredentialError::RandomSource)?;
Ok(generation)
}
fn new_bootstrap(generation: [u8; 16]) -> Result<CodexBootstrapCredential, CredentialError> {
let mut secret = [0u8; TOKEN_BYTES];
getrandom::getrandom(&mut secret).map_err(|_| CredentialError::RandomSource)?;
Ok(CodexBootstrapCredential {
digest: credential_digest(&secret),
secret,
generation,
revoked: false,
})
}
#[cfg(unix)]
fn write_bootstrap_secret(
runtime_dir: PathBuf,
secret: &[u8; TOKEN_BYTES],
generation: [u8; 16],
) -> Result<CodexBootstrapFile, CredentialError> {
use std::io::Write;
use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt};
let metadata = std::fs::symlink_metadata(&runtime_dir)?;
if metadata.file_type().is_symlink()
|| !metadata.is_dir()
|| metadata.permissions().mode() & 0o777 != 0o700
|| metadata.uid() != unsafe { libc::geteuid() }
{
return Err(CredentialError::Io(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"runtime credential directory must be owner-owned mode 0700 and not a symlink",
)));
}
for _ in 0..3 {
let mut name_bytes = [0u8; 16];
getrandom::getrandom(&mut name_bytes).map_err(|_| CredentialError::RandomSource)?;
let name = name_bytes
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>();
name_bytes.zeroize();
let path = runtime_dir.join(format!("bootstrap-{name}.credential"));
let mut options = std::fs::OpenOptions::new();
options
.write(true)
.create_new(true)
.mode(0o600)
.custom_flags(libc::O_NOFOLLOW | libc::O_CLOEXEC);
let mut file = match options.open(&path) {
Ok(file) => file,
Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => continue,
Err(error) => return Err(CredentialError::Io(error)),
};
let mut encoded = encode_hex(secret);
let write_result = file
.write_all(encoded.as_bytes())
.and_then(|_| file.sync_all());
encoded.zeroize();
if let Err(error) = write_result {
let _ = std::fs::remove_file(&path);
return Err(CredentialError::Io(error));
}
let written = match file.metadata() {
Ok(metadata) => metadata,
Err(error) => {
drop(file);
let _ = std::fs::remove_file(&path);
return Err(CredentialError::Io(error));
}
};
if !written.is_file()
|| written.permissions().mode() & 0o777 != 0o600
|| written.uid() != unsafe { libc::geteuid() }
{
drop(file);
let _ = std::fs::remove_file(&path);
return Err(CredentialError::Io(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"bootstrap credential file failed owner/mode verification",
)));
}
return Ok(CodexBootstrapFile { path, generation });
}
Err(CredentialError::Collision)
}
fn encode_hex(secret: &[u8; TOKEN_BYTES]) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut output = String::with_capacity(TOKEN_BYTES * 2);
for byte in secret {
output.push(HEX[(byte >> 4) as usize] as char);
output.push(HEX[(byte & 0x0f) as usize] as char);
}
output
}
fn decode_hex_secret(value: &str) -> Option<[u8; TOKEN_BYTES]> {
if value.len() != TOKEN_BYTES * 2 {
return None;
}
let mut output = [0u8; TOKEN_BYTES];
for (index, pair) in value.as_bytes().chunks_exact(2).enumerate() {
output[index] = (hex_nibble(pair[0])? << 4) | hex_nibble(pair[1])?;
}
Some(output)
}
fn hex_nibble(value: u8) -> Option<u8> {
match value {
b'0'..=b'9' => Some(value - b'0'),
b'a'..=b'f' => Some(value - b'a' + 10),
_ => None,
}
}
fn zero(secret: &mut [u8; TOKEN_BYTES]) {
secret.zeroize();
}