use crate::cancellation::{AgentCancellation, AgentCancellationHandle};
use anyhow::{Context, Result, bail};
use std::{
io::{Read, Write},
net::{TcpListener, TcpStream},
sync::{Arc, Mutex},
thread::{self, JoinHandle},
time::{Duration, Instant},
};
const MAX_HEADERS: usize = 32 * 1024;
const MAX_REQUEST: usize = 8 * 1024 * 1024;
const MAX_RESPONSE: usize = 16 * 1024 * 1024;
const SOCKET_POLL: Duration = Duration::from_millis(100);
pub(super) const CLAUDE_STARTUP_TIMEOUT: Duration = Duration::from_secs(300);
pub(super) const CLAUDE_MAX_REQUEST_DURATION: Duration = Duration::from_secs(30 * 60);
#[derive(Default)]
struct AdmissionState {
claimed: bool,
denied: u32,
result: Option<RelayResult>,
failed: bool,
phase: &'static str,
}
pub(super) struct RelayResult {
pub(super) status: u16,
pub(super) body: Vec<u8>,
pub(super) denied: u32,
}
pub(super) struct ClaudeAdmission {
pub(super) url: String,
state: Arc<Mutex<AdmissionState>>,
stop: AgentCancellationHandle,
cancellation: AgentCancellation,
worker: Option<JoinHandle<()>>,
}
impl ClaudeAdmission {
pub(super) fn start(
cancellation: &AgentCancellation,
host_messages: Vec<serde_json::Value>,
) -> Result<Self> {
Self::start_at(
"https://api.anthropic.com".to_string(),
cancellation,
host_messages,
)
}
fn start_at(
upstream: String,
cancellation: &AgentCancellation,
host_messages: Vec<serde_json::Value>,
) -> Result<Self> {
if upstream != "https://api.anthropic.com" {
bail!("Claude subscription relay requires official Anthropic endpoint");
}
let listener = TcpListener::bind(("127.0.0.1", 0)).context("binding Claude relay")?;
listener.set_nonblocking(true)?;
let route = format!("/admit/{}", uuid::Uuid::new_v4());
let url = format!("http://127.0.0.1:{}{route}", listener.local_addr()?.port());
let (worker_cancel, stop) = cancellation.child_token();
let state = Arc::new(Mutex::new(AdmissionState::default()));
let worker_state = Arc::clone(&state);
let worker = thread::spawn(move || {
while !worker_cancel.is_canceled() {
match listener.accept() {
Ok((socket, _)) => {
if worker_cancel.is_canceled() {
break;
}
if serve(
socket,
&route,
&upstream,
&host_messages,
&worker_state,
&worker_cancel,
)
.is_err()
&& !worker_cancel.is_canceled()
{
worker_state
.lock()
.unwrap_or_else(|p| p.into_inner())
.failed = true;
}
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(SOCKET_POLL);
}
Err(_) => {
worker_state
.lock()
.unwrap_or_else(|p| p.into_inner())
.failed = true;
break;
}
}
}
});
Ok(Self {
url,
state,
stop,
cancellation: cancellation.clone(),
worker: Some(worker),
})
}
pub(super) fn finish(mut self) -> Result<RelayResult> {
let deadline = Instant::now() + Duration::from_secs(5);
while Instant::now() < deadline {
self.cancellation.check()?;
let state = self.state.lock().unwrap_or_else(|p| p.into_inner());
if state.failed || state.result.is_some() {
break;
}
drop(state);
thread::sleep(SOCKET_POLL);
}
self.cancellation.check()?;
self.shutdown();
let mut state = self.state.lock().unwrap_or_else(|p| p.into_inner());
if let Some(result) = state.result.as_ref().filter(|result| result.status != 200) {
return Err(upstream_error_diagnostic(result.status, &result.body));
}
if state.denied != 0 {
bail!(
"Claude CLI attempted {} additional model requests; magi-code admits one request per generation (claimed={}, failed={}, phase={})",
state.denied,
state.claimed,
state.failed,
state.phase
);
}
if !state.claimed || state.failed {
bail!(
"Claude relay incomplete (claimed={}, failed={}, phase={})",
state.claimed,
state.failed,
state.phase
);
}
let mut result = state
.result
.take()
.context("Claude relay missing response")?;
result.denied = state.denied;
Ok(result)
}
fn shutdown(&mut self) {
self.stop.cancel();
if let Some(worker) = self.worker.take() {
let _ = worker.join();
}
}
}
impl Drop for ClaudeAdmission {
fn drop(&mut self) {
self.shutdown();
}
}
const PROMPT_TOO_LONG_HINT: &str = "Claude CLI rejected prompt as too long; run /compact or start /new; session history was not truncated";
pub(super) fn claude_error_with_hint(error: anyhow::Error, hint: Option<&str>) -> anyhow::Error {
match hint {
Some(PROMPT_TOO_LONG_HINT) if !crate::cancellation::is_run_canceled(&error) => {
super::error::ProviderError::prompt_too_long(format!("{error}; {PROMPT_TOO_LONG_HINT}"))
.into()
}
Some(hint) => {
let message = format!("{error}; {hint}");
error.context(message)
}
None => error,
}
}
pub(super) fn claude_error_hint(message: &[u8]) -> Option<&'static str> {
let text = String::from_utf8_lossy(&message[..message.len().min(8192)]).to_ascii_lowercase();
if text.contains("prompt") && text.contains("too long") {
Some(PROMPT_TOO_LONG_HINT)
} else if text.contains("model")
&& [
"not available",
"does not exist",
"not found",
"not allowed",
"no access",
"don't have access",
"do not have access",
"not authorized",
"invalid model",
"unknown model",
]
.iter()
.any(|phrase| text.contains(phrase))
{
Some("model unavailable for this Claude account or CLI; choose a supported model")
} else if text.contains("rate_limit") || text.contains("rate limit") {
Some("Claude rate limit reached; retry later")
} else if text.contains("authentication_error") || text.contains("unauthorized") {
Some("Claude authentication rejected; check claude auth status")
} else {
None
}
}
fn upstream_error_diagnostic(status: u16, body: &[u8]) -> anyhow::Error {
let hint = (body.len() <= 8192)
.then(|| serde_json::from_slice::<serde_json::Value>(body).ok())
.flatten()
.and_then(|error| {
error
.pointer("/error/message")
.and_then(serde_json::Value::as_str)
.and_then(|message| claude_error_hint(message.as_bytes()))
})
.or(match status {
401 | 403 => {
Some("Claude authentication or model access rejected; check claude auth status")
}
429 => Some("Claude rate limit reached; retry later"),
_ => None,
});
claude_error_with_hint(
anyhow::anyhow!("Claude upstream returned HTTP {status}"),
hint,
)
}
fn restore_host_conversation(
payload: &[u8],
host_messages: &[serde_json::Value],
) -> Option<Vec<u8>> {
use serde_json::{Value, json};
const MAX_CACHE_BREAKPOINTS: usize = 4;
if host_messages.is_empty() {
return None;
}
let mut body: Value = serde_json::from_slice(payload).ok()?;
let cli_messages = body.get("messages")?.as_array()?;
let mut reminders = Vec::new();
let mut cli_system_messages = Vec::new();
for message in cli_messages {
match message.get("role").and_then(Value::as_str)? {
"system" => cli_system_messages.push(message.clone()),
"user" => reminders.extend(
message
.get("content")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter(|block| {
block
.get("text")
.and_then(Value::as_str)
.is_some_and(|text| text.starts_with("<system-reminder>"))
})
.cloned(),
),
_ => {}
}
}
let mut messages = host_messages.to_vec();
let used_breakpoints = serde_json::to_string(&body["system"])
.ok()?
.matches("\"cache_control\"")
.count()
+ serde_json::to_string(&body["tools"])
.ok()?
.matches("\"cache_control\"")
.count()
+ serde_json::to_string(&cli_system_messages)
.ok()?
.matches("\"cache_control\"")
.count();
if used_breakpoints < MAX_CACHE_BREAKPOINTS
&& let Some(last_block) = messages
.last_mut()
.and_then(|message| message.get_mut("content"))
.and_then(Value::as_array_mut)
.and_then(|content| content.last_mut())
{
last_block["cache_control"] = json!({"type":"ephemeral","ttl":"1h"});
}
if !reminders.is_empty() {
let first_user_content = messages
.iter_mut()
.find(|message| message["role"] == "user")
.and_then(|message| message.get_mut("content"))
.and_then(Value::as_array_mut)?;
first_user_content.splice(0..0, reminders);
}
messages.extend(cli_system_messages);
body["messages"] = Value::Array(messages);
serde_json::to_vec(&body).ok()
}
fn serve(
mut socket: TcpStream,
route: &str,
upstream: &str,
host_messages: &[serde_json::Value],
state: &Mutex<AdmissionState>,
cancel: &AgentCancellation,
) -> Result<()> {
socket.set_read_timeout(Some(SOCKET_POLL))?;
socket.set_write_timeout(Some(SOCKET_POLL))?;
let deadline = Instant::now() + CLAUDE_STARTUP_TIMEOUT;
let mut data = Vec::new();
let mut buffer = [0u8; 8192];
let end = loop {
cancel.check()?;
if Instant::now() >= deadline {
bail!("Claude relay request timed out");
}
if data.len() > MAX_HEADERS {
bail!("Claude relay request headers too large");
}
match socket.read(&mut buffer) {
Ok(0) => bail!("Claude relay request ended early"),
Ok(count) => data.extend_from_slice(&buffer[..count]),
Err(error)
if matches!(
error.kind(),
std::io::ErrorKind::TimedOut | std::io::ErrorKind::WouldBlock
) =>
{
continue;
}
Err(_) => bail!("Claude relay request read failed"),
}
if let Some(offset) = data.windows(4).position(|part| part == b"\r\n\r\n") {
break offset + 4;
}
};
if end > MAX_HEADERS {
bail!("Claude relay request headers too large");
}
let header_text = std::str::from_utf8(&data[..end]).context("invalid Claude relay headers")?;
let mut lines = header_text.split("\r\n");
let requested = lines.next().context("missing Claude relay request line")?;
let Some(path) = requested
.strip_prefix("POST ")
.and_then(|rest| rest.strip_suffix(" HTTP/1.1"))
else {
write_response(&mut socket, 404, b"", cancel)?;
return Ok(());
};
let expected = format!("{route}/v1/messages");
if path.split('?').next() != Some(expected.as_str()) {
write_response(&mut socket, 404, b"", cancel)?;
return Ok(());
}
let mut headers = reqwest::header::HeaderMap::new();
let mut length = None;
for line in lines.filter(|line| !line.is_empty()) {
let (key, value) = line
.split_once(':')
.context("invalid Claude relay header")?;
let key = reqwest::header::HeaderName::from_bytes(key.as_bytes())
.context("invalid Claude relay header name")?;
let value = reqwest::header::HeaderValue::from_str(value.trim())
.context("invalid Claude relay header value")?;
if key == reqwest::header::CONTENT_LENGTH {
if length.is_some() {
bail!("duplicate Claude relay content length");
}
length = Some(
value
.to_str()?
.parse::<usize>()
.context("invalid Claude relay content length")?,
);
}
if key == reqwest::header::ORIGIN || key == reqwest::header::TRANSFER_ENCODING {
bail!("unsupported Claude relay request header");
}
if key.as_str() != "proxy-connection"
&& !matches!(
key,
reqwest::header::HOST
| reqwest::header::CONNECTION
| reqwest::header::CONTENT_LENGTH
| reqwest::header::ACCEPT_ENCODING
| reqwest::header::PROXY_AUTHORIZATION
)
{
headers.append(key, value);
}
}
{
let mut state = state.lock().unwrap_or_else(|p| p.into_inner());
if state.claimed {
state.denied = state.denied.saturating_add(1);
drop(state);
write_response(&mut socket, 400, br#"{"type":"error","error":{"type":"invalid_request_error","message":"MODEL_ADMISSION_CONSUMED"}}"#, cancel)?;
return Ok(());
}
state.claimed = true;
state.phase = "reading request body";
}
let length = length.context("missing Claude relay content length")?;
if length > MAX_REQUEST {
bail!("Claude relay request too large");
}
let mut payload = data[end..].to_vec();
if payload.len() > length {
bail!("Claude relay request exceeds content length");
}
while payload.len() < length {
cancel.check()?;
if Instant::now() >= deadline {
bail!("Claude relay request timed out");
}
let count = buffer.len().min(length - payload.len());
match socket.read(&mut buffer[..count]) {
Ok(0) => bail!("Claude relay request body ended early"),
Ok(count) => payload.extend_from_slice(&buffer[..count]),
Err(error)
if matches!(
error.kind(),
std::io::ErrorKind::TimedOut | std::io::ErrorKind::WouldBlock
) =>
{
continue;
}
Err(_) => bail!("Claude relay request body read failed"),
}
}
let payload = restore_host_conversation(&payload, host_messages).unwrap_or(payload);
state.lock().unwrap_or_else(|p| p.into_inner()).phase = "connecting upstream";
headers.insert(
reqwest::header::ACCEPT_ENCODING,
reqwest::header::HeaderValue::from_static("identity"),
);
let query = path
.split_once('?')
.map(|(_, value)| format!("?{value}"))
.unwrap_or_default();
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()?;
let canceled = async {
while !cancel.is_canceled() {
tokio::time::sleep(SOCKET_POLL).await;
}
};
let relay = runtime.block_on(async {
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(CLAUDE_STARTUP_TIMEOUT)
.timeout(CLAUDE_MAX_REQUEST_DURATION)
.build()?;
let request = client
.post(format!("{upstream}/v1/messages{query}"))
.headers(headers)
.body(payload);
let upstream = async {
let mut response = request.send().await?;
state.lock().unwrap_or_else(|p| p.into_inner()).phase = "streaming upstream";
let status = response.status().as_u16();
let mut body = Vec::new();
let mut socket_started = false;
while let Some(chunk) = response.chunk().await? {
if body.len().saturating_add(chunk.len()) > MAX_RESPONSE {
bail!("Claude upstream response too large");
}
if !socket_started {
write_stream_headers(&mut socket, status, cancel)?;
socket_started = true;
}
body.extend_from_slice(&chunk);
write_chunk(&mut socket, &chunk, cancel)?;
}
if !socket_started {
write_stream_headers(&mut socket, status, cancel)?;
}
write_all_cancellable(&mut socket, b"0\r\n\r\n", cancel)?;
state.lock().unwrap_or_else(|p| p.into_inner()).phase = "response complete";
Ok::<_, anyhow::Error>(RelayResult {
status,
body,
denied: 0,
})
};
tokio::select! {
result = upstream => result,
_ = canceled => bail!("Claude relay canceled"),
}
});
state.lock().unwrap_or_else(|p| p.into_inner()).result = Some(relay?);
Ok(())
}
fn write_all_cancellable(
socket: &mut TcpStream,
mut bytes: &[u8],
cancel: &AgentCancellation,
) -> Result<()> {
while !bytes.is_empty() {
cancel.check()?;
match socket.write(bytes) {
Ok(0) => bail!("Claude relay socket closed"),
Ok(count) => bytes = &bytes[count..],
Err(error)
if matches!(
error.kind(),
std::io::ErrorKind::TimedOut | std::io::ErrorKind::WouldBlock
) =>
{
continue;
}
Err(_) => bail!("Claude relay socket write failed"),
}
}
Ok(())
}
fn write_stream_headers(
socket: &mut TcpStream,
status: u16,
cancel: &AgentCancellation,
) -> Result<()> {
let content_type = if status == 200 {
"text/event-stream"
} else {
"application/json"
};
write_all_cancellable(socket, format!("HTTP/1.1 {status} Upstream\r\nTransfer-Encoding: chunked\r\nContent-Type: {content_type}\r\nConnection: close\r\n\r\n").as_bytes(), cancel)
}
fn write_chunk(socket: &mut TcpStream, chunk: &[u8], cancel: &AgentCancellation) -> Result<()> {
write_all_cancellable(socket, format!("{:X}\r\n", chunk.len()).as_bytes(), cancel)?;
write_all_cancellable(socket, chunk, cancel)?;
write_all_cancellable(socket, b"\r\n", cancel)
}
fn write_response(
socket: &mut TcpStream,
status: u16,
body: &[u8],
cancel: &AgentCancellation,
) -> Result<()> {
write_all_cancellable(socket, format!("HTTP/1.1 {status} Denied\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n", body.len()).as_bytes(), cancel)?;
write_all_cancellable(socket, body, cancel)
}