1mod config;
4
5use std::{
6 collections::{HashMap, VecDeque},
7 io::Read as _,
8 path::{Path, PathBuf},
9 sync::{
10 Arc,
11 atomic::{AtomicU64, AtomicUsize, Ordering},
12 },
13 time::Duration,
14};
15
16use anyhow::{Context, Result, anyhow};
17use async_trait::async_trait;
18use config::Config;
19pub use config::{ApprovalPolicy, ConfigOverrides};
20pub fn init_user_config() -> anyhow::Result<std::path::PathBuf> { config::Config::init_user_config() }
21use scv_core::{
22 AgentError, AgentRuntime, ApprovalGate, ApprovalRequest, BudgetContextPolicy, CoreEvent,
23 EventSink, Message, ToolRegistry, ToolRisk,
24};
25use scv_protocol::{ClientMessage, PROTOCOL_VERSION, PeerInfo, QueueEntry, ServerEvent, Usage};
26use scv_provider_openai::OpenAiProvider;
27use scv_tools::{SkillMap, builtin_registry};
28use tokio::{
29 io::{AsyncBufRead, AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader},
30 net::{UnixListener, UnixStream},
31 sync::{Mutex, OwnedSemaphorePermit, Semaphore, mpsc, oneshot},
32 task::JoinHandle,
33};
34use tokio_util::sync::CancellationToken;
35use uuid::Uuid;
36
37const PROMPT_LIMIT_BYTES: usize = 256 * 1024;
38const OUTPUT_QUEUE_CAPACITY: usize = 256;
39const OUTPUT_QUEUE_MIN_BYTES: usize = 16 * 1024 * 1024;
40const SHUTDOWN_GRACE: Duration = Duration::from_secs(3);
41const MAX_QUEUE_ITEMS: usize = 64;
42const MAX_QUEUE_BYTES: usize = 4 * 1024 * 1024;
43
44pub async fn run_stdio(overrides: ConfigOverrides) -> Result<()> {
45 let stdin = tokio::io::stdin();
46 let stdout = tokio::io::stdout();
47 run(stdin, stdout, overrides).await
48}
49
50pub fn default_socket_path() -> Result<PathBuf> {
52 let root = std::env::var_os("SCV_HOME")
53 .map(PathBuf::from)
54 .or_else(|| dirs::home_dir().map(|path| path.join(".scv")))
55 .ok_or_else(|| anyhow!("cannot determine SCV_HOME"))?;
56 Ok(root.join("server.sock"))
57}
58
59pub async fn run_socket(path: &Path, overrides: ConfigOverrides) -> Result<()> {
61 if let Some(parent) = path.parent() {
62 tokio::fs::create_dir_all(parent).await.context("create SCV socket directory")?;
63 }
64 if path.exists() {
65 if UnixStream::connect(path).await.is_ok() {
66 return Err(anyhow!("SCV server is already running at {}", path.display()));
67 }
68 tokio::fs::remove_file(path).await.with_context(|| format!("remove stale SCV socket {}", path.display()))?;
69 }
70 let listener = UnixListener::bind(path).with_context(|| format!("bind SCV server socket {}", path.display()))?;
71 #[cfg(unix)] {
72 use std::os::unix::fs::PermissionsExt;
73 std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600)).context("secure SCV socket")?;
74 }
75 let result = loop {
76 tokio::select! {
77 accepted = listener.accept() => {
78 let (stream, _) = accepted.context("accept SCV client")?;
79 let child_overrides = overrides.clone();
80 tokio::spawn(async move {
81 let (reader, writer) = stream.into_split();
82 if let Err(error) = run(reader, writer, child_overrides).await {
83 tracing::warn!(%error, "SCV socket client stopped");
84 }
85 });
86 }
87 _ = tokio::signal::ctrl_c() => break Ok(()),
88 }
89 };
90 drop(listener);
91 let _ = tokio::fs::remove_file(path).await;
92 result
93}
94
95async fn run<R, W>(reader: R, writer: W, overrides: ConfigOverrides) -> Result<()>
96where
97 R: tokio::io::AsyncRead + Unpin,
98 W: tokio::io::AsyncWrite + Unpin + Send + 'static,
99{
100 let initial_output_bytes =
101 output_queue_bytes(Config::default().protocol.max_server_frame_bytes)?;
102 let (output_tx, mut output_rx) = outbound_channel(initial_output_bytes);
103 let mut writer_task = tokio::spawn(async move {
104 let mut writer = writer;
105 while let Some(frame) = output_rx.recv().await {
106 writer.write_all(&frame.bytes).await?;
107 writer.write_all(b"\n").await?;
108 writer.flush().await?;
109 }
110 Ok::<(), std::io::Error>(())
111 });
112 let (done_tx, mut done_rx) = mpsc::channel::<TurnDone>(4);
113 let approvals = Arc::new(ApprovalBroker::default());
114 let mut reader = BufReader::new(reader);
115 let mut initialized = false;
116 let mut session: Option<Session> = None;
117 let mut active: Option<ActiveTurn> = None;
118 let mut fatal = false;
119 let mut writer_finished = false;
120
121 let loop_result: Result<()> = async {
122 loop {
123 let frame_limit = session.as_ref().map_or_else(
124 || Config::default().protocol.max_client_frame_bytes,
125 |value| value.config.protocol.max_client_frame_bytes,
126 );
127 tokio::select! {
128 read = read_bounded_frame(&mut reader, frame_limit) => {
129 let frame = match read.context("read protocol input")? {
130 FrameRead::Eof => {
131 if let Some(active) = &active { active.cancellation.cancel(); }
132 break;
133 }
134 FrameRead::TooLarge => {
135 send_error(&output_tx, "", "invalid_request", "client frame exceeds configured limit", false, server_frame_limit(&session)).await?;
136 continue;
137 }
138 FrameRead::Frame(frame) => frame,
139 };
140 if frame.is_empty() {
141 send_error(&output_tx, "", "invalid_json", "protocol frame is empty", false, server_frame_limit(&session)).await?;
142 continue;
143 }
144 let message = match serde_json::from_slice::<ClientMessage>(&frame) {
145 Ok(message) => message,
146 Err(error) => {
147 send_error(&output_tx, "", "invalid_json", &format!("invalid protocol JSON: {error}"), false, server_frame_limit(&session)).await?;
148 continue;
149 }
150 };
151 match message {
152 ClientMessage::Initialize { request_id, protocol_version, .. } => {
153 if initialized {
154 send_error(&output_tx, &request_id, "invalid_request", "connection is already initialized", false, server_frame_limit(&session)).await?;
155 continue;
156 }
157 if protocol_version != PROTOCOL_VERSION {
158 send_error(&output_tx, &request_id, "version_mismatch", &format!("server supports protocol {PROTOCOL_VERSION}"), true, server_frame_limit(&session)).await?;
159 fatal = true;
160 break;
161 }
162 initialized = true;
163 send_event(&output_tx, ServerEvent::Initialized {
164 request_id,
165 protocol_version: PROTOCOL_VERSION,
166 server: PeerInfo { name: "scv-server".into(), version: env!("CARGO_PKG_VERSION").into() },
167 }, Config::default().protocol.max_server_frame_bytes).await?;
168 }
169 other if !initialized => {
170 send_error(&output_tx, other.request_id(), "not_initialized", "initialize must be the first message", false, server_frame_limit(&session)).await?;
171 }
172 ClientMessage::SessionStart { request_id, cwd, provider, model, base_url, no_tools } => {
173 if session.is_some() {
174 send_error(&output_tx, &request_id, "invalid_request", "this connection already has a session", false, server_frame_limit(&session)).await?;
175 continue;
176 }
177 let session_overrides = ConfigOverrides {
178 provider: provider.or_else(|| overrides.provider.clone()),
179 model: model.or_else(|| overrides.model.clone()),
180 base_url: base_url.or_else(|| overrides.base_url.clone()),
181 approval_policy: overrides.approval_policy,
182 no_tools: no_tools.unwrap_or(overrides.no_tools),
183 };
184 match build_session(&cwd, session_overrides).await {
185 Ok(new_session) => {
186 output_tx.ensure_capacity(output_queue_bytes(
187 new_session.config.protocol.max_server_frame_bytes,
188 )?)?;
189 let event = ServerEvent::SessionStarted {
190 request_id,
191 session_id: new_session.id.clone(),
192 cwd: new_session.workspace.display().to_string(),
193 model: new_session.runtime.model().to_owned(),
194 context_max_tokens: new_session.config.context.max_tokens,
195 max_server_frame_bytes: new_session.config.protocol.max_server_frame_bytes,
196 max_transcript_bytes: new_session.config.tui.max_transcript_bytes,
197 max_transcript_items: new_session.config.tui.max_transcript_items,
198 max_prompt_history_bytes: new_session.config.tui.max_prompt_history_bytes,
199 max_prompt_history_items: new_session.config.tui.max_prompt_history_items,
200 };
201 send_event(&output_tx, event, new_session.config.protocol.max_server_frame_bytes).await?;
202 send_event(&output_tx, ServerEvent::QueueSnapshot {
203 request_id: None,
204 session_id: new_session.id.clone(),
205 seq: next_seq(&new_session.seq),
206 entries: new_session.queue.lock().await.iter().cloned().collect(),
207 paused: new_session.paused.load(Ordering::Acquire),
208 }, new_session.config.protocol.max_server_frame_bytes).await?;
209 session = Some(new_session);
210 }
211 Err(error) => {
212 send_error(&output_tx, &request_id, "invalid_request", &error.to_string(), false, server_frame_limit(&session)).await?;
213 }
214 }
215 }
216 ClientMessage::SessionAttach { request_id, .. } => {
217 send_error(&output_tx, &request_id, "unsupported", "session attach requires the shared socket server", false, server_frame_limit(&session)).await?;
218 }
219 ClientMessage::TurnStart { request_id, session_id, prompt } => {
220 let Some(current) = session.as_ref() else {
221 send_error(&output_tx, &request_id, "session_not_found", "start a session first", false, server_frame_limit(&session)).await?;
222 continue;
223 };
224 if current.id != session_id {
225 send_error(&output_tx, &request_id, "session_not_found", "session id does not match", false, server_frame_limit(&session)).await?;
226 continue;
227 }
228 if prompt.trim().is_empty() || prompt.len() > PROMPT_LIMIT_BYTES {
229 send_error(&output_tx, &request_id, "invalid_request", "prompt must be non-empty and no larger than 256 KiB", false, server_frame_limit(&session)).await?;
230 continue;
231 }
232 if active.is_some() {
233 let entry = match current.enqueue(prompt, request_id.clone()).await {
234 Ok(entry) => entry,
235 Err(code) => { send_error(&output_tx, &request_id, code, "session queue limit reached", false, server_frame_limit(&session)).await?; continue; }
236 };
237 let position = current.queue.lock().await.len().saturating_sub(1);
238 send_event(&output_tx, ServerEvent::QueueEnqueued {
239 request_id, session_id: current.id.clone(), seq: next_seq(¤t.seq), entry, position,
240 }, current.config.protocol.max_server_frame_bytes).await?;
241 continue;
242 }
243 let turn_id = Uuid::new_v4().to_string();
244 let cancellation = CancellationToken::new();
245 let meta = TurnMeta {
246 request_id: request_id.clone(),
247 session_id: current.id.clone(),
248 turn_id: turn_id.clone(),
249 seq: Arc::clone(¤t.seq),
250 max_server_frame: current.config.protocol.max_server_frame_bytes,
251 };
252 send_event(&output_tx, ServerEvent::TurnStarted {
253 request_id: request_id.clone(),
254 session_id: current.id.clone(),
255 turn_id: turn_id.clone(),
256 seq: next_seq(¤t.seq),
257 }, current.config.protocol.max_server_frame_bytes).await?;
258 let runtime = Arc::clone(¤t.runtime);
259 let history = Arc::clone(¤t.history);
260 let sink: Arc<dyn EventSink> = Arc::new(ProtocolSink {
261 meta: meta.clone(),
262 output: output_tx.clone(),
263 cancellation: cancellation.clone(),
264 });
265 let gate: Arc<dyn ApprovalGate> = Arc::new(ProtocolApprovalGate {
266 policy: current.config.tools.approval_policy,
267 broker: Arc::clone(&approvals),
268 meta,
269 output: output_tx.clone(),
270 });
271 let task_cancel = cancellation.clone();
272 let task_done = done_tx.clone();
273 let task_request = request_id.clone();
274 let task_session = current.id.clone();
275 let task_turn = turn_id.clone();
276 let task = tokio::spawn(async move {
277 let mut history = history.lock().await;
278 let result = runtime.run_turn(&mut history, prompt, sink, gate, task_cancel).await;
279 let _ = task_done.send(TurnDone {
280 request_id: task_request,
281 session_id: task_session,
282 turn_id: task_turn,
283 result,
284 }).await;
285 });
286 active = Some(ActiveTurn { turn_id, cancellation, task });
287 }
288 ClientMessage::QueueUpdate { request_id, session_id, queue_id, revision, prompt } => {
289 let Some(current) = session.as_ref() else { send_error(&output_tx, &request_id, "session_not_found", "start a session first", false, server_frame_limit(&session)).await?; continue; };
290 if current.id != session_id { send_error(&output_tx, &request_id, "session_not_found", "session id does not match", false, server_frame_limit(&session)).await?; continue; }
291 if prompt.trim().is_empty() || prompt.len() > PROMPT_LIMIT_BYTES { send_error(&output_tx, &request_id, "invalid_request", "prompt must be non-empty and no larger than 256 KiB", false, server_frame_limit(&session)).await?; continue; }
292 match current.update_queue(&queue_id, revision, prompt).await {
293 Ok(entry) => send_event(&output_tx, ServerEvent::QueueUpdated { request_id, session_id: current.id.clone(), seq: next_seq(¤t.seq), entry }, current.config.protocol.max_server_frame_bytes).await?,
294 Err(code) => send_error(&output_tx, &request_id, &code, "queue entry was not found or revision is stale", false, server_frame_limit(&session)).await?,
295 }
296 }
297 ClientMessage::QueueMove { request_id, session_id, queue_id, revision, before_queue_id } => {
298 let Some(current) = session.as_ref() else { send_error(&output_tx, &request_id, "session_not_found", "start a session first", false, server_frame_limit(&session)).await?; continue; };
299 match current.move_queue(&session_id, &queue_id, revision, before_queue_id).await {
300 Ok((id, rev, pos)) => send_event(&output_tx, ServerEvent::QueueMoved { request_id, session_id: current.id.clone(), seq: next_seq(¤t.seq), queue_id: id, position: pos, revision: rev }, current.config.protocol.max_server_frame_bytes).await?,
301 Err(code) => send_error(&output_tx, &request_id, &code, "queue entry was not found or revision is stale", false, server_frame_limit(&session)).await?,
302 }
303 }
304 ClientMessage::QueueRemove { request_id, session_id, queue_id, revision } => {
305 let Some(current) = session.as_ref() else { send_error(&output_tx, &request_id, "session_not_found", "start a session first", false, server_frame_limit(&session)).await?; continue; };
306 match current.remove_queue(&session_id, &queue_id, revision).await {
307 Ok((id, rev)) => send_event(&output_tx, ServerEvent::QueueRemoved { request_id, session_id: current.id.clone(), seq: next_seq(¤t.seq), queue_id: id, revision: rev }, current.config.protocol.max_server_frame_bytes).await?,
308 Err(code) => send_error(&output_tx, &request_id, &code, "queue entry was not found or revision is stale", false, server_frame_limit(&session)).await?,
309 }
310 }
311 ClientMessage::SessionPause { request_id, session_id, paused } => {
312 let Some(current) = session.as_ref() else { send_error(&output_tx, &request_id, "session_not_found", "start a session first", false, server_frame_limit(&session)).await?; continue; };
313 if current.id != session_id { send_error(&output_tx, &request_id, "session_not_found", "session id does not match", false, server_frame_limit(&session)).await?; continue; }
314 current.paused.store(paused, Ordering::Release);
315 send_event(&output_tx, ServerEvent::SessionPaused { request_id, session_id: current.id.clone(), seq: next_seq(¤t.seq), paused }, current.config.protocol.max_server_frame_bytes).await?;
316 }
317 ClientMessage::TurnCancel { request_id, session_id, turn_id } => {
318 match (&session, &active) {
319 (Some(current), Some(running)) if current.id == session_id && running.turn_id == turn_id => running.cancellation.cancel(),
320 _ => send_error(&output_tx, &request_id, "turn_not_found", "active turn was not found", false, server_frame_limit(&session)).await?,
321 }
322 }
323 ClientMessage::ApprovalResolve { request_id, session_id, approval_id, approved } => {
324 if session.as_ref().is_none_or(|current| current.id != session_id) {
325 send_error(&output_tx, &request_id, "session_not_found", "session id does not match", false, server_frame_limit(&session)).await?;
326 } else if !approvals.resolve(&approval_id, approved).await {
327 send_error(&output_tx, &request_id, "approval_not_found", "approval was not found or already resolved", false, server_frame_limit(&session)).await?;
328 }
329 }
330 ClientMessage::SessionClear { request_id, session_id } => {
331 let Some(current) = session.as_ref() else {
332 send_error(&output_tx, &request_id, "session_not_found", "session was not found", false, server_frame_limit(&session)).await?;
333 continue;
334 };
335 if current.id != session_id {
336 send_error(&output_tx, &request_id, "session_not_found", "session id does not match", false, server_frame_limit(&session)).await?;
337 } else if active.is_some() {
338 send_error(&output_tx, &request_id, "turn_active", "cancel the active turn before clearing", false, server_frame_limit(&session)).await?;
339 } else {
340 current.history.lock().await.clear();
341 current.queue.lock().await.clear();
342 send_event(&output_tx, ServerEvent::SessionCleared {
343 request_id,
344 session_id: current.id.clone(),
345 seq: next_seq(¤t.seq),
346 }, current.config.protocol.max_server_frame_bytes).await?;
347 send_event(&output_tx, ServerEvent::QueueSnapshot { request_id: None, session_id: current.id.clone(), seq: next_seq(¤t.seq), entries: Vec::new(), paused: current.paused.load(Ordering::Acquire) }, current.config.protocol.max_server_frame_bytes).await?;
348 }
349 }
350 }
351 }
352 writer = &mut writer_task => {
353 writer_finished = true;
354 writer.context("join protocol writer")??;
355 break;
356 }
357 done = done_rx.recv(), if active.is_some() => {
358 if let Some(done) = done {
359 if let Some(current) = session.as_ref() {
360 let seq = next_seq(¤t.seq);
361 let event = match done.result {
362 Ok(outcome) => ServerEvent::TurnCompleted {
363 request_id: done.request_id,
364 session_id: done.session_id,
365 turn_id: done.turn_id,
366 seq,
367 steps: outcome.steps,
368 usage: Usage { input_tokens: outcome.usage.input_tokens, output_tokens: outcome.usage.output_tokens },
369 },
370 Err(AgentError::Cancelled) => ServerEvent::TurnCancelled {
371 request_id: done.request_id,
372 session_id: done.session_id,
373 turn_id: done.turn_id,
374 seq,
375 },
376 Err(error) => ServerEvent::TurnFailed {
377 request_id: done.request_id,
378 session_id: done.session_id,
379 turn_id: done.turn_id,
380 seq,
381 code: error.code().into(),
382 message: error.to_string(),
383 },
384 };
385 send_event(&output_tx, event, current.config.protocol.max_server_frame_bytes).await?;
386 }
387 if let Some(active) = active.take() {
388 let _ = active.task.await;
389 }
390 if let Some(current) = session.as_ref()
391 && !current.paused.load(Ordering::Acquire)
392 && let Some(entry) = current.queue.lock().await.pop_front()
393 {
394 let turn_id = Uuid::new_v4().to_string();
395 let cancellation = CancellationToken::new();
396 send_event(&output_tx, ServerEvent::QueueDequeued {
397 request_id: entry.submitter.clone(),
398 session_id: current.id.clone(),
399 seq: next_seq(¤t.seq),
400 queue_id: entry.queue_id,
401 turn_id: turn_id.clone(),
402 }, current.config.protocol.max_server_frame_bytes).await?;
403 send_event(&output_tx, ServerEvent::TurnStarted {
404 request_id: entry.submitter.clone(),
405 session_id: current.id.clone(),
406 turn_id: turn_id.clone(),
407 seq: next_seq(¤t.seq),
408 }, current.config.protocol.max_server_frame_bytes).await?;
409 let meta = TurnMeta {
410 request_id: entry.submitter.clone(), session_id: current.id.clone(), turn_id: turn_id.clone(),
411 seq: Arc::clone(¤t.seq), max_server_frame: current.config.protocol.max_server_frame_bytes,
412 };
413 let sink: Arc<dyn EventSink> = Arc::new(ProtocolSink { meta: meta.clone(), output: output_tx.clone(), cancellation: cancellation.clone() });
414 let gate: Arc<dyn ApprovalGate> = Arc::new(ProtocolApprovalGate { policy: current.config.tools.approval_policy, broker: Arc::clone(&approvals), meta, output: output_tx.clone() });
415 let runtime = Arc::clone(¤t.runtime);
416 let history = Arc::clone(¤t.history);
417 let task_done = done_tx.clone();
418 let task_request = entry.submitter;
419 let task_session = current.id.clone();
420 let task_turn = turn_id.clone();
421 let task_cancel = cancellation.clone();
422 let task = tokio::spawn(async move {
423 let mut history = history.lock().await;
424 let result = runtime.run_turn(&mut history, entry.prompt, sink, gate, task_cancel).await;
425 let _ = task_done.send(TurnDone { request_id: task_request, session_id: task_session, turn_id: task_turn, result }).await;
426 });
427 active = Some(ActiveTurn { turn_id, cancellation, task });
428 }
429 }
430 }
431 }
432 }
433 Ok(())
434 }
435 .await;
436
437 if let Some(active) = active.take() {
438 shutdown_active_turn(active, SHUTDOWN_GRACE).await;
439 }
440 drop(output_tx);
441 let writer_result = if writer_finished {
442 Ok(())
443 } else {
444 shutdown_writer(writer_task, SHUTDOWN_GRACE).await
445 };
446 loop_result?;
447 writer_result?;
448 if fatal {
449 return Err(anyhow!("protocol version mismatch"));
450 }
451 Ok(())
452}
453
454enum FrameRead {
455 Eof,
456 Frame(Vec<u8>),
457 TooLarge,
458}
459
460struct OutboundFrame {
461 bytes: Vec<u8>,
462 _byte_permit: OwnedSemaphorePermit,
463}
464
465#[derive(Clone)]
466struct OutboundSender {
467 frames: mpsc::Sender<OutboundFrame>,
468 budget: Arc<Semaphore>,
469 capacity: Arc<AtomicUsize>,
470}
471
472#[derive(Debug, PartialEq, Eq)]
473enum OutboundSendError {
474 Cancelled,
475 Closed,
476 TimedOut,
477 FrameExceedsQueue { frame_bytes: usize, capacity: usize },
478}
479
480impl std::fmt::Display for OutboundSendError {
481 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
482 match self {
483 Self::Cancelled => formatter.write_str("outbound send cancelled"),
484 Self::Closed => formatter.write_str("protocol client disconnected"),
485 Self::TimedOut => formatter.write_str("outbound send timed out under backpressure"),
486 Self::FrameExceedsQueue {
487 frame_bytes,
488 capacity,
489 } => write!(
490 formatter,
491 "outbound frame uses {frame_bytes} bytes but queue capacity is {capacity} bytes"
492 ),
493 }
494 }
495}
496
497impl std::error::Error for OutboundSendError {}
498
499fn outbound_channel(capacity: usize) -> (OutboundSender, mpsc::Receiver<OutboundFrame>) {
500 let (frames, receiver) = mpsc::channel(OUTPUT_QUEUE_CAPACITY);
501 (
502 OutboundSender {
503 frames,
504 budget: Arc::new(Semaphore::new(capacity)),
505 capacity: Arc::new(AtomicUsize::new(capacity)),
506 },
507 receiver,
508 )
509}
510
511impl OutboundSender {
512 fn ensure_capacity(&self, required: usize) -> Result<()> {
513 if required > Semaphore::MAX_PERMITS {
514 return Err(anyhow!(
515 "outbound queue capacity {required} exceeds runtime limit {}",
516 Semaphore::MAX_PERMITS
517 ));
518 }
519 let current = self.capacity.load(Ordering::Acquire);
520 if required > current {
521 self.budget.add_permits(required - current);
522 self.capacity.store(required, Ordering::Release);
523 }
524 Ok(())
525 }
526
527 async fn send(
528 &self,
529 bytes: Vec<u8>,
530 cancellation: Option<&CancellationToken>,
531 ) -> std::result::Result<(), OutboundSendError> {
532 self.send_with_timeout(bytes, cancellation, SHUTDOWN_GRACE)
533 .await
534 }
535
536 async fn send_with_timeout(
537 &self,
538 bytes: Vec<u8>,
539 cancellation: Option<&CancellationToken>,
540 control_timeout: Duration,
541 ) -> std::result::Result<(), OutboundSendError> {
542 let frame_bytes =
543 bytes
544 .len()
545 .checked_add(1)
546 .ok_or(OutboundSendError::FrameExceedsQueue {
547 frame_bytes: usize::MAX,
548 capacity: self.capacity.load(Ordering::Acquire),
549 })?;
550 let capacity = self.capacity.load(Ordering::Acquire);
551 let permits =
552 u32::try_from(frame_bytes).map_err(|_| OutboundSendError::FrameExceedsQueue {
553 frame_bytes,
554 capacity,
555 })?;
556 if frame_bytes > capacity {
557 return Err(OutboundSendError::FrameExceedsQueue {
558 frame_bytes,
559 capacity,
560 });
561 }
562
563 let control_deadline = tokio::time::Instant::now() + control_timeout;
564 let acquire = Arc::clone(&self.budget).acquire_many_owned(permits);
565 let permit = if let Some(cancellation) = cancellation {
566 tokio::select! {
567 biased;
568 _ = cancellation.cancelled() => return Err(OutboundSendError::Cancelled),
569 permit = acquire => permit.map_err(|_| OutboundSendError::Closed)?,
570 }
571 } else {
572 tokio::time::timeout_at(control_deadline, acquire)
573 .await
574 .map_err(|_| OutboundSendError::TimedOut)?
575 .map_err(|_| OutboundSendError::Closed)?
576 };
577 let frame = OutboundFrame {
578 bytes,
579 _byte_permit: permit,
580 };
581 if let Some(cancellation) = cancellation {
582 tokio::select! {
583 biased;
584 _ = cancellation.cancelled() => Err(OutboundSendError::Cancelled),
585 result = self.frames.send(frame) => result.map_err(|_| OutboundSendError::Closed),
586 }
587 } else {
588 tokio::time::timeout_at(control_deadline, self.frames.send(frame))
589 .await
590 .map_err(|_| OutboundSendError::TimedOut)?
591 .map_err(|_| OutboundSendError::Closed)
592 }
593 }
594}
595
596fn output_queue_bytes(max_frame_bytes: usize) -> Result<usize> {
597 let required = max_frame_bytes
598 .checked_add(1)
599 .and_then(|bytes| bytes.checked_mul(2))
600 .ok_or_else(|| anyhow!("configured server frame limit is too large"))?
601 .max(OUTPUT_QUEUE_MIN_BYTES);
602 if required > Semaphore::MAX_PERMITS {
603 return Err(anyhow!(
604 "configured server frame limit requires an outbound queue larger than the runtime supports"
605 ));
606 }
607 Ok(required)
608}
609
610async fn read_bounded_frame<R>(reader: &mut R, max_bytes: usize) -> std::io::Result<FrameRead>
611where
612 R: AsyncBufRead + Unpin,
613{
614 let mut frame = Vec::with_capacity(max_bytes.min(8192));
615 let read_limit = u64::try_from(max_bytes)
616 .unwrap_or(u64::MAX)
617 .saturating_add(2);
618 let mut limited = reader.take(read_limit);
619 let read = limited.read_until(b'\n', &mut frame).await?;
620 drop(limited);
621 if read == 0 {
622 return Ok(FrameRead::Eof);
623 }
624 let ended_with_newline = frame.last() == Some(&b'\n');
625 while matches!(frame.last(), Some(b'\n' | b'\r')) {
626 frame.pop();
627 }
628 if frame.len() <= max_bytes {
629 return Ok(FrameRead::Frame(frame));
630 }
631 if !ended_with_newline {
632 loop {
633 let available = reader.fill_buf().await?;
634 if available.is_empty() {
635 break;
636 }
637 if let Some(end) = available.iter().position(|byte| *byte == b'\n') {
638 reader.consume(end + 1);
639 break;
640 }
641 let consumed = available.len();
642 reader.consume(consumed);
643 }
644 }
645 Ok(FrameRead::TooLarge)
646}
647
648fn server_frame_limit(session: &Option<Session>) -> usize {
649 session.as_ref().map_or_else(
650 || Config::default().protocol.max_server_frame_bytes,
651 |value| value.config.protocol.max_server_frame_bytes,
652 )
653}
654
655struct Session {
656 id: String,
657 workspace: PathBuf,
658 config: Config,
659 runtime: Arc<AgentRuntime>,
660 history: Arc<Mutex<Vec<Message>>>,
661 seq: Arc<AtomicU64>,
662 queue: Arc<Mutex<VecDeque<QueueEntry>>>,
663 paused: Arc<std::sync::atomic::AtomicBool>,
664}
665
666impl Session {
667 async fn enqueue(&self, prompt: String, submitter: String) -> std::result::Result<QueueEntry, &'static str> {
668 let entry = QueueEntry { queue_id: Uuid::new_v4().to_string(), revision: 1, prompt, submitter };
669 let mut queue = self.queue.lock().await;
670 let bytes: usize = queue.iter().map(|item| item.prompt.len()).sum();
671 if queue.len() >= MAX_QUEUE_ITEMS || bytes.saturating_add(entry.prompt.len()) > MAX_QUEUE_BYTES {
672 return Err("queue_limit");
673 }
674 queue.push_back(entry.clone());
675 Ok(entry)
676 }
677
678 async fn update_queue(&self, id: &str, revision: u64, prompt: String) -> std::result::Result<QueueEntry, &'static str> {
679 let mut queue = self.queue.lock().await;
680 let bytes: usize = queue.iter().map(|item| item.prompt.len()).sum();
681 let entry = queue.iter_mut().find(|entry| entry.queue_id == id).ok_or("queue_not_found")?;
682 if entry.revision != revision { return Err("queue_conflict"); }
683 if bytes.saturating_sub(entry.prompt.len()).saturating_add(prompt.len()) > MAX_QUEUE_BYTES { return Err("queue_limit"); }
684 entry.prompt = prompt;
685 entry.revision += 1;
686 Ok(entry.clone())
687 }
688
689 async fn move_queue(&self, session_id: &str, id: &str, revision: u64, before: Option<String>) -> std::result::Result<(String, u64, usize), &'static str> {
690 if self.id != session_id { return Err("session_not_found"); }
691 let mut queue = self.queue.lock().await;
692 let index = queue.iter().position(|entry| entry.queue_id == id).ok_or("queue_not_found")?;
693 if queue[index].revision != revision { return Err("queue_conflict"); }
694 let target_index = match before.as_deref() {
697 Some(target) if target == id => return Ok((id.to_string(), revision, index)),
698 Some(target) => Some(queue.iter().position(|item| item.queue_id == target).ok_or("queue_not_found")?),
699 None => None,
700 };
701 let mut entry = queue.remove(index).expect("queue index exists");
702 let target = target_index.map_or(queue.len(), |target| target.saturating_sub(usize::from(target > index)));
703 let pos = target.min(queue.len());
704 let id = entry.queue_id.clone();
705 let rev = entry.revision + 1;
706 entry.revision = rev;
707 queue.insert(pos, entry);
708 Ok((id, rev, pos))
709 }
710
711 async fn remove_queue(&self, session_id: &str, id: &str, revision: u64) -> std::result::Result<(String, u64), &'static str> {
712 if self.id != session_id { return Err("session_not_found"); }
713 let mut queue = self.queue.lock().await;
714 let index = queue.iter().position(|entry| entry.queue_id == id).ok_or("queue_not_found")?;
715 if queue[index].revision != revision { return Err("queue_conflict"); }
716 let entry = queue.remove(index).expect("queue index exists");
717 Ok((entry.queue_id, entry.revision))
718 }
719}
720
721struct ActiveTurn {
722 turn_id: String,
723 cancellation: CancellationToken,
724 task: JoinHandle<()>,
725}
726
727async fn shutdown_active_turn(mut active: ActiveTurn, grace: Duration) -> bool {
728 active.cancellation.cancel();
729 if tokio::time::timeout(grace, &mut active.task).await.is_ok() {
730 true
731 } else {
732 active.task.abort();
733 let _ = active.task.await;
734 false
735 }
736}
737
738async fn shutdown_writer(
739 mut writer: JoinHandle<std::io::Result<()>>,
740 grace: Duration,
741) -> Result<()> {
742 match tokio::time::timeout(grace, &mut writer).await {
743 Ok(result) => {
744 result.context("join protocol writer")??;
745 Ok(())
746 }
747 Err(_) => {
748 writer.abort();
749 let _ = writer.await;
750 Err(anyhow!("protocol writer shutdown timed out"))
751 }
752 }
753}
754
755struct TurnDone {
756 request_id: String,
757 session_id: String,
758 turn_id: String,
759 result: Result<scv_core::TurnOutcome, AgentError>,
760}
761
762#[derive(Clone)]
763struct TurnMeta {
764 request_id: String,
765 session_id: String,
766 turn_id: String,
767 seq: Arc<AtomicU64>,
768 max_server_frame: usize,
769}
770
771fn next_seq(sequence: &AtomicU64) -> u64 {
772 sequence.fetch_add(1, Ordering::Relaxed) + 1
773}
774
775async fn build_session(cwd: &str, overrides: ConfigOverrides) -> Result<Session> {
776 let workspace = std::fs::canonicalize(cwd).with_context(|| format!("resolve cwd {cwd}"))?;
777 if !workspace.is_dir() {
778 return Err(anyhow!("cwd is not a directory"));
779 }
780 let no_tools = overrides.no_tools;
781 let config = Config::load(&workspace, overrides)?;
782 let provider_config = config.provider.clone();
783 let api_key = provider_config.api_key.clone().or_else(|| {
784 provider_config.api_key_env.as_deref().and_then(|name| std::env::var(name).ok())
785 }).filter(|key| !key.trim().is_empty()).ok_or_else(|| anyhow!("provider credential is not configured; set provider.api_key or provider.api_key_env"))?;
786 let (skills, skill_roots, skill_prompt) = discover_skills(&workspace, &config)?;
787 let system_prompt = build_system_prompt(&workspace, &config, &skill_prompt)?;
788 let provider = Arc::new(OpenAiProvider::new(
789 provider_config.model.clone(),
790 provider_config.base_url.clone(),
791 api_key,
792 Duration::from_secs(provider_config.timeout_seconds),
793 config.provider_limits(),
794 provider_config.headers.clone(),
795 )?);
796 let tools = if no_tools {
797 Arc::new(ToolRegistry::default())
798 } else {
799 Arc::new(builtin_registry(
800 config.tools(),
801 skills,
802 skill_roots,
803 config.skills.max_skill_bytes,
804 config.adapters(),
805 )?)
806 };
807 let context = Arc::new(BudgetContextPolicy::new((&config.context).into())?);
808 let runtime = Arc::new(AgentRuntime::new(
809 provider,
810 tools,
811 context,
812 config.core_agent(system_prompt),
813 workspace.clone(),
814 ));
815 Ok(Session {
816 id: Uuid::new_v4().to_string(),
817 workspace,
818 config,
819 runtime,
820 history: Arc::new(Mutex::new(Vec::new())),
821 seq: Arc::new(AtomicU64::new(0)),
822 queue: Arc::new(Mutex::new(VecDeque::new())),
823 paused: Arc::new(std::sync::atomic::AtomicBool::new(false)),
824 })
825}
826
827fn build_system_prompt(workspace: &Path, config: &Config, skills: &str) -> Result<String> {
828 let mut prompt = config.agent.system_prompt.clone();
829 prompt.push_str(&format!(
830 "\nCurrent working directory: {}\n",
831 workspace.display()
832 ));
833 let agents_path = workspace.join("AGENTS.md");
834 if agents_path.is_file() {
835 let canonical = std::fs::canonicalize(&agents_path).context("resolve project AGENTS.md")?;
836 if !canonical.starts_with(workspace) {
837 return Err(anyhow!("project AGENTS.md escaped workspace"));
838 }
839 let (bytes, truncated) = read_prefix(&canonical, config.tools.max_read_bytes)
840 .context("read project AGENTS.md")?;
841 let instructions = std::str::from_utf8(&bytes).context("project AGENTS.md is not UTF-8")?;
842 prompt.push_str("\n# Project instructions\n");
843 prompt.push_str(instructions);
844 if truncated {
845 prompt.push_str("\n[AGENTS.md truncated by configured read limit]\n");
846 }
847 }
848 if !skills.is_empty() {
849 prompt.push_str("\n# Available skills\n");
850 prompt.push_str(skills);
851 prompt.push_str("\nUse read_skill with a skill name when its workflow applies.\n");
852 }
853 Ok(prompt)
854}
855
856fn discover_skills(workspace: &Path, config: &Config) -> Result<(SkillMap, Vec<PathBuf>, String)> {
857 let mut skills = SkillMap::new();
858 let mut roots = Vec::new();
859 let project_root = workspace.join(&config.skills.project_dir);
860 for (root, must_be_workspace) in [(&project_root, true), (&config.skills.user_dir, false)] {
861 if !root.is_dir() {
862 continue;
863 }
864 let canonical = std::fs::canonicalize(root)
865 .with_context(|| format!("resolve skill root {}", root.display()))?;
866 if must_be_workspace && !canonical.starts_with(workspace) {
867 return Err(anyhow!("project skill root escaped workspace"));
868 }
869 roots.push(canonical.clone());
870 let mut entries: Vec<_> = std::fs::read_dir(&canonical)
871 .with_context(|| format!("read skill root {}", canonical.display()))?
872 .filter_map(Result::ok)
873 .collect();
874 entries.sort_by_key(|entry| entry.file_name());
875 for entry in entries {
876 if skills.len() >= config.skills.max_skills {
877 break;
878 }
879 let path = entry.path().join("SKILL.md");
880 if !path.is_file() {
881 continue;
882 }
883 let canonical_file = std::fs::canonicalize(&path)
884 .with_context(|| format!("resolve skill {}", path.display()))?;
885 if !canonical_file.starts_with(&canonical) {
886 continue;
887 }
888 let name = entry.file_name().to_string_lossy().to_string();
889 skills.entry(name).or_insert(canonical_file);
890 }
891 }
892 let mut names: Vec<_> = skills.keys().cloned().collect();
893 names.sort();
894 let mut listing = String::new();
895 for name in names {
896 let path = &skills[&name];
897 let bytes = read_prefix(path, config.skills.max_skill_bytes)
898 .map(|(bytes, _)| bytes)
899 .unwrap_or_default();
900 let content = String::from_utf8_lossy(&bytes);
901 let description = skill_description(&content);
902 listing.push_str(&format!("- {name}: {description}\n"));
903 }
904 Ok((skills, roots, listing))
905}
906
907fn read_prefix(path: &Path, max_bytes: usize) -> std::io::Result<(Vec<u8>, bool)> {
908 let file = std::fs::File::open(path)?;
909 let mut bytes = Vec::with_capacity(max_bytes.min(8192));
910 file.take(
911 u64::try_from(max_bytes)
912 .unwrap_or(u64::MAX)
913 .saturating_add(1),
914 )
915 .read_to_end(&mut bytes)?;
916 let truncated = bytes.len() > max_bytes;
917 bytes.truncate(max_bytes);
918 Ok((bytes, truncated))
919}
920
921fn skill_description(content: &str) -> String {
922 if let Some(frontmatter) = content.strip_prefix("---\n")
923 && let Some((header, _)) = frontmatter.split_once("\n---")
924 {
925 for line in header.lines() {
926 if let Some(description) = line.strip_prefix("description:") {
927 return description.trim().trim_matches('"').to_owned();
928 }
929 }
930 }
931 content
932 .lines()
933 .map(str::trim)
934 .find(|line| !line.is_empty() && !line.starts_with('#'))
935 .unwrap_or("No description provided")
936 .chars()
937 .take(240)
938 .collect()
939}
940
941struct ProtocolSink {
942 meta: TurnMeta,
943 output: OutboundSender,
944 cancellation: CancellationToken,
945}
946
947#[async_trait]
948impl EventSink for ProtocolSink {
949 async fn emit(&self, event: CoreEvent) -> Result<(), AgentError> {
950 let seq = next_seq(&self.meta.seq);
951 let event = match event {
952 CoreEvent::AssistantDelta { content } => ServerEvent::AssistantDelta {
953 request_id: self.meta.request_id.clone(),
954 session_id: self.meta.session_id.clone(),
955 turn_id: self.meta.turn_id.clone(),
956 seq,
957 content,
958 },
959 CoreEvent::AssistantCompleted { content } => ServerEvent::AssistantCompleted {
960 request_id: self.meta.request_id.clone(),
961 session_id: self.meta.session_id.clone(),
962 turn_id: self.meta.turn_id.clone(),
963 seq,
964 content,
965 },
966 CoreEvent::ToolProposed {
967 call_id,
968 name,
969 arguments,
970 } => ServerEvent::ToolProposed {
971 request_id: self.meta.request_id.clone(),
972 session_id: self.meta.session_id.clone(),
973 turn_id: self.meta.turn_id.clone(),
974 seq,
975 call_id,
976 name,
977 arguments,
978 },
979 CoreEvent::ToolStarted { call_id, name } => ServerEvent::ToolStarted {
980 request_id: self.meta.request_id.clone(),
981 session_id: self.meta.session_id.clone(),
982 turn_id: self.meta.turn_id.clone(),
983 seq,
984 call_id,
985 name,
986 },
987 CoreEvent::ToolCompleted {
988 call_id,
989 name,
990 output,
991 } => ServerEvent::ToolCompleted {
992 request_id: self.meta.request_id.clone(),
993 session_id: self.meta.session_id.clone(),
994 turn_id: self.meta.turn_id.clone(),
995 seq,
996 call_id,
997 name,
998 success: !output.is_error,
999 output: output.content,
1000 truncated: output.truncated,
1001 },
1002 CoreEvent::ContextCompacted {
1003 before_tokens,
1004 after_tokens,
1005 removed_messages,
1006 } => ServerEvent::ContextCompacted {
1007 request_id: self.meta.request_id.clone(),
1008 session_id: self.meta.session_id.clone(),
1009 turn_id: self.meta.turn_id.clone(),
1010 seq,
1011 before_tokens,
1012 after_tokens,
1013 removed_messages,
1014 },
1015 CoreEvent::SessionTrimmed {
1016 removed_messages,
1017 history_bytes,
1018 } => ServerEvent::SessionTrimmed {
1019 request_id: self.meta.request_id.clone(),
1020 session_id: self.meta.session_id.clone(),
1021 seq,
1022 removed_messages,
1023 history_bytes,
1024 },
1025 };
1026 send_turn_event(
1027 &self.output,
1028 event,
1029 self.meta.max_server_frame,
1030 &self.cancellation,
1031 )
1032 .await
1033 }
1034}
1035
1036#[derive(Default)]
1037struct ApprovalBroker {
1038 pending: Mutex<HashMap<String, oneshot::Sender<bool>>>,
1039}
1040
1041impl ApprovalBroker {
1042 async fn insert(&self, id: String, sender: oneshot::Sender<bool>) {
1043 self.pending.lock().await.insert(id, sender);
1044 }
1045
1046 async fn remove(&self, id: &str) {
1047 self.pending.lock().await.remove(id);
1048 }
1049
1050 async fn resolve(&self, id: &str, approved: bool) -> bool {
1051 let sender = self.pending.lock().await.remove(id);
1052 sender.is_some_and(|sender| sender.send(approved).is_ok())
1053 }
1054}
1055
1056struct ProtocolApprovalGate {
1057 policy: ApprovalPolicy,
1058 broker: Arc<ApprovalBroker>,
1059 meta: TurnMeta,
1060 output: OutboundSender,
1061}
1062
1063#[async_trait]
1064impl ApprovalGate for ProtocolApprovalGate {
1065 async fn approve(
1066 &self,
1067 request: ApprovalRequest,
1068 cancellation: CancellationToken,
1069 ) -> Result<bool, AgentError> {
1070 match self.policy {
1071 ApprovalPolicy::OnRisk if request.risk == ToolRisk::ReadOnly => return Ok(true),
1072 ApprovalPolicy::Never => return Ok(request.risk == ToolRisk::ReadOnly),
1073 ApprovalPolicy::Always | ApprovalPolicy::OnRisk => {}
1074 }
1075 let approval_id = Uuid::new_v4().to_string();
1076 let (sender, receiver) = oneshot::channel();
1077 self.broker.insert(approval_id.clone(), sender).await;
1078 let event = ServerEvent::ApprovalRequested {
1079 request_id: self.meta.request_id.clone(),
1080 session_id: self.meta.session_id.clone(),
1081 turn_id: self.meta.turn_id.clone(),
1082 seq: next_seq(&self.meta.seq),
1083 approval_id: approval_id.clone(),
1084 call_id: request.call_id,
1085 name: request.name,
1086 risk: request.risk.as_str().into(),
1087 cwd: request.cwd.display().to_string(),
1088 summary: request.summary,
1089 };
1090 if let Err(error) = send_turn_event(
1091 &self.output,
1092 event,
1093 self.meta.max_server_frame,
1094 &cancellation,
1095 )
1096 .await
1097 {
1098 self.broker.remove(&approval_id).await;
1099 return Err(error);
1100 }
1101 tokio::select! {
1102 result = receiver => result.map_err(|_| AgentError::Cancelled),
1103 _ = cancellation.cancelled() => {
1104 self.broker.remove(&approval_id).await;
1105 Err(AgentError::Cancelled)
1106 }
1107 }
1108 }
1109}
1110
1111async fn send_event(output: &OutboundSender, event: ServerEvent, max_bytes: usize) -> Result<()> {
1112 let bytes = encode_event(&event, max_bytes)?;
1113 output.send(bytes, None).await.map_err(anyhow::Error::new)
1114}
1115
1116async fn send_turn_event(
1117 output: &OutboundSender,
1118 event: ServerEvent,
1119 max_bytes: usize,
1120 cancellation: &CancellationToken,
1121) -> Result<(), AgentError> {
1122 let bytes = encode_event(&event, max_bytes)
1123 .map_err(|error| AgentError::ResponseLimit(error.to_string()))?;
1124 match output.send(bytes, Some(cancellation)).await {
1125 Ok(()) => Ok(()),
1126 Err(OutboundSendError::Cancelled) => Err(AgentError::Cancelled),
1127 Err(error) => Err(AgentError::Internal(error.to_string())),
1128 }
1129}
1130
1131fn encode_event(event: &ServerEvent, max_bytes: usize) -> Result<Vec<u8>> {
1132 let bytes = serde_json::to_vec(event).context("serialize protocol event")?;
1133 if bytes.len() > max_bytes {
1134 return Err(anyhow!("server event exceeds configured frame limit"));
1135 }
1136 Ok(bytes)
1137}
1138
1139async fn send_error(
1140 output: &OutboundSender,
1141 request_id: &str,
1142 code: &str,
1143 message: &str,
1144 fatal: bool,
1145 max_bytes: usize,
1146) -> Result<()> {
1147 send_event(
1148 output,
1149 ServerEvent::Error {
1150 request_id: (!request_id.is_empty()).then(|| request_id.to_owned()),
1151 code: code.into(),
1152 message: message.into(),
1153 fatal,
1154 },
1155 max_bytes,
1156 )
1157 .await
1158}
1159
1160#[cfg(test)]
1161mod tests {
1162 use std::{
1163 future::pending,
1164 io::Cursor,
1165 sync::atomic::{AtomicBool, Ordering},
1166 };
1167
1168 use super::*;
1169
1170 struct DropSignal(Arc<AtomicBool>);
1171
1172 impl Drop for DropSignal {
1173 fn drop(&mut self) {
1174 self.0.store(true, Ordering::Release);
1175 }
1176 }
1177
1178 #[tokio::test]
1179 async fn bounded_reader_discards_an_oversized_line() {
1180 let input = format!("{}\n{{}}\n", "x".repeat(10));
1181 let mut reader = BufReader::new(Cursor::new(input.into_bytes()));
1182 assert!(matches!(
1183 read_bounded_frame(&mut reader, 4).await.unwrap(),
1184 FrameRead::TooLarge
1185 ));
1186 match read_bounded_frame(&mut reader, 4).await.unwrap() {
1187 FrameRead::Frame(frame) => assert_eq!(frame, b"{}"),
1188 _ => panic!("expected the frame following the oversized line"),
1189 }
1190 }
1191
1192 #[tokio::test]
1193 async fn bounded_reader_accepts_exact_crlf_limit() {
1194 let mut reader = BufReader::new(Cursor::new(b"1234\r\n".to_vec()));
1195 match read_bounded_frame(&mut reader, 4).await.unwrap() {
1196 FrameRead::Frame(frame) => assert_eq!(frame, b"1234"),
1197 _ => panic!("expected an exact-limit frame"),
1198 }
1199 }
1200
1201 #[tokio::test]
1202 async fn outbound_byte_backpressure_is_cancellation_aware() {
1203 let (output, mut receiver) = outbound_channel(5);
1204 output.send(vec![0; 4], None).await.unwrap();
1205
1206 let cancellation = CancellationToken::new();
1207 let blocked = tokio::spawn({
1208 let output = output.clone();
1209 let cancellation = cancellation.clone();
1210 async move { output.send(vec![1; 4], Some(&cancellation)).await }
1211 });
1212 tokio::task::yield_now().await;
1213 assert!(!blocked.is_finished());
1214
1215 cancellation.cancel();
1216 assert_eq!(blocked.await.unwrap(), Err(OutboundSendError::Cancelled));
1217
1218 drop(receiver.recv().await.unwrap());
1219 output.send(vec![2; 4], None).await.unwrap();
1220 }
1221
1222 #[tokio::test]
1223 async fn outbound_control_send_times_out_under_byte_backpressure() {
1224 let (output, _receiver) = outbound_channel(5);
1225 output.send(vec![0; 4], None).await.unwrap();
1226 let result = output
1227 .send_with_timeout(vec![1; 4], None, Duration::from_millis(10))
1228 .await;
1229 assert_eq!(result, Err(OutboundSendError::TimedOut));
1230 }
1231
1232 #[tokio::test]
1233 async fn active_turn_shutdown_aborts_after_grace_period() {
1234 let cancellation = CancellationToken::new();
1235 let dropped = Arc::new(AtomicBool::new(false));
1236 let (started_tx, started_rx) = oneshot::channel();
1237 let task = tokio::spawn({
1238 let dropped = Arc::clone(&dropped);
1239 async move {
1240 let _signal = DropSignal(dropped);
1241 let _ = started_tx.send(());
1242 pending::<()>().await;
1243 }
1244 });
1245 started_rx.await.unwrap();
1246
1247 let graceful = shutdown_active_turn(
1248 ActiveTurn {
1249 turn_id: "turn".into(),
1250 cancellation,
1251 task,
1252 },
1253 Duration::from_millis(10),
1254 )
1255 .await;
1256
1257 assert!(!graceful);
1258 assert!(dropped.load(Ordering::Acquire));
1259 }
1260
1261 #[tokio::test]
1262 async fn writer_shutdown_aborts_after_grace_period() {
1263 let dropped = Arc::new(AtomicBool::new(false));
1264 let (started_tx, started_rx) = oneshot::channel();
1265 let writer = tokio::spawn({
1266 let dropped = Arc::clone(&dropped);
1267 async move {
1268 let _signal = DropSignal(dropped);
1269 let _ = started_tx.send(());
1270 pending::<std::io::Result<()>>().await
1271 }
1272 });
1273 started_rx.await.unwrap();
1274
1275 let result = shutdown_writer(writer, Duration::from_millis(10)).await;
1276
1277 assert!(result.is_err());
1278 assert!(dropped.load(Ordering::Acquire));
1279 }
1280}