Skip to main content

scv_server/
lib.rs

1//! SCV's authoritative stdio server.
2
3mod 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(&current.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(&current.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(&current.seq),
204                        }, current.config.protocol.max_server_frame_bytes).await?;
205                        let runtime = Arc::clone(&current.runtime);
206                        let history = Arc::clone(&current.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(&current.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(&current.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(&current.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(&current.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(&current.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(&current.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(&current.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(&current.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(&current.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(&current.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(&current.runtime);
363                        let history = Arc::clone(&current.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        // Validate the destination while the source is still present. This keeps
642        // the operation atomic and handles a self move as a no-op reorder.
643        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}