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