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