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