github_copilot_sdk/
subscription.rs1use std::collections::VecDeque;
36use std::fmt;
37use std::pin::Pin;
38use std::sync::Arc;
39use std::task::{Context, Poll};
40
41use parking_lot::Mutex;
42use tokio::sync::broadcast::{Receiver, Sender, WeakSender};
43use tokio_stream::wrappers::BroadcastStream;
44use tokio_stream::wrappers::errors::BroadcastStreamRecvError;
45use tokio_stream::{Stream, StreamExt as _};
46
47use crate::types::{SessionEvent, SessionLifecycleEvent};
48use crate::{Custom, Repr};
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq)]
58pub struct Lagged(pub(crate) u64);
59
60impl Lagged {
61 pub fn skipped(&self) -> u64 {
63 self.0
64 }
65}
66
67impl fmt::Display for Lagged {
68 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
69 write!(f, "subscription lagged behind by {} events", self.0)
70 }
71}
72
73impl std::error::Error for Lagged {}
74
75#[derive(Clone, Copy, Debug, PartialEq, Eq)]
77#[non_exhaustive]
78pub enum RecvErrorKind {
79 Closed,
82
83 Lagged(Lagged),
85}
86
87impl fmt::Display for RecvErrorKind {
88 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
89 match self {
90 RecvErrorKind::Closed => write!(f, "subscription closed"),
91 RecvErrorKind::Lagged(l) => write!(f, "{l}"),
92 }
93 }
94}
95
96#[derive(Debug)]
99pub struct RecvError {
100 repr: Repr<RecvErrorKind>,
101}
102
103impl RecvError {
104 pub fn kind(&self) -> &RecvErrorKind {
106 match &self.repr {
107 Repr::Simple(k) | Repr::SimpleMessage(k, ..) | Repr::Custom(Custom { kind: k, .. }) => {
108 k
109 }
110 }
111 }
112}
113
114impl fmt::Display for RecvError {
115 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
116 match &self.repr {
117 Repr::Simple(k) => write!(f, "{k}"),
118 Repr::SimpleMessage(_, m) => write!(f, "{m}"),
119 Repr::Custom(Custom { error, .. }) => write!(f, "{error}"),
120 }
121 }
122}
123
124impl std::error::Error for RecvError {
125 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
126 match &self.repr {
127 Repr::Custom(Custom { error, .. }) => Some(&**error),
128 _ => None,
129 }
130 }
131}
132
133impl From<RecvErrorKind> for RecvError {
134 fn from(kind: RecvErrorKind) -> Self {
135 Self {
136 repr: Repr::Simple(kind),
137 }
138 }
139}
140
141impl From<Lagged> for RecvError {
142 fn from(lagged: Lagged) -> Self {
143 Self::from(RecvErrorKind::Lagged(lagged))
144 }
145}
146
147enum ResumeBootstrapState {
148 Unclaimed(VecDeque<SessionEvent>),
149 Claimed(VecDeque<SessionEvent>),
150 Disabled,
151}
152
153pub(crate) struct ResumeBootstrap {
156 state: Mutex<ResumeBootstrapState>,
157 live: WeakSender<SessionEvent>,
158}
159
160pub(crate) struct ResumeBootstrapCleanup(Arc<ResumeBootstrap>);
162
163impl Drop for ResumeBootstrapCleanup {
164 fn drop(&mut self) {
165 self.0.release_unclaimed();
166 }
167}
168
169impl ResumeBootstrap {
170 pub(crate) fn new(event_tx: &Sender<SessionEvent>) -> Arc<Self> {
171 Arc::new(Self {
172 state: Mutex::new(ResumeBootstrapState::Unclaimed(VecDeque::new())),
173 live: event_tx.downgrade(),
174 })
175 }
176
177 pub(crate) fn cleanup_guard(self: &Arc<Self>) -> ResumeBootstrapCleanup {
178 ResumeBootstrapCleanup(self.clone())
179 }
180
181 pub(crate) fn publish(&self, event_tx: &Sender<SessionEvent>, event: SessionEvent) {
182 let mut state = self.state.lock();
183 match &mut *state {
184 ResumeBootstrapState::Unclaimed(events) | ResumeBootstrapState::Claimed(events) => {
185 events.push_back(event.clone());
186 }
187 ResumeBootstrapState::Disabled => {}
188 }
189 let _ = event_tx.send(event);
191 }
192
193 pub(crate) fn subscribe(
194 self: &Arc<Self>,
195 event_tx: &Sender<SessionEvent>,
196 ) -> EventSubscription {
197 let mut state = self.state.lock();
198 match &mut *state {
199 ResumeBootstrapState::Unclaimed(events) => {
200 let events = std::mem::take(events);
201 *state = ResumeBootstrapState::Claimed(events);
202 EventSubscription {
203 inner: None,
204 bootstrap: Some(self.clone()),
205 }
206 }
207 ResumeBootstrapState::Claimed(_) | ResumeBootstrapState::Disabled => {
208 EventSubscription::new(event_tx.subscribe())
209 }
210 }
211 }
212
213 fn pop(&self, live: &mut Option<BroadcastStream<SessionEvent>>) -> Option<SessionEvent> {
214 let mut state = self.state.lock();
215 let ResumeBootstrapState::Claimed(events) = &mut *state else {
216 return None;
217 };
218 if let Some(event) = events.pop_front() {
219 return Some(event);
220 }
221 *live = self
224 .live
225 .upgrade()
226 .map(|sender| BroadcastStream::new(sender.subscribe()));
227 *state = ResumeBootstrapState::Disabled;
228 None
229 }
230
231 pub(crate) fn release_unclaimed(&self) {
232 let mut state = self.state.lock();
233 if matches!(*state, ResumeBootstrapState::Unclaimed(_)) {
234 *state = ResumeBootstrapState::Disabled;
235 }
236 }
237
238 fn abandon(&self) {
239 let mut state = self.state.lock();
240 if matches!(*state, ResumeBootstrapState::Claimed(_)) {
241 *state = ResumeBootstrapState::Disabled;
242 }
243 }
244}
245
246#[must_use = "dropping the subscription unsubscribes and discards any owned resume bootstrap backlog"]
255pub struct EventSubscription {
256 inner: Option<BroadcastStream<SessionEvent>>,
257 bootstrap: Option<Arc<ResumeBootstrap>>,
258}
259
260impl EventSubscription {
261 pub(crate) fn new(rx: Receiver<SessionEvent>) -> Self {
262 Self {
263 inner: Some(BroadcastStream::new(rx)),
264 bootstrap: None,
265 }
266 }
267
268 fn next_bootstrap_event(&mut self) -> Option<SessionEvent> {
269 let event = self
270 .bootstrap
271 .as_ref()
272 .and_then(|bootstrap| bootstrap.pop(&mut self.inner));
273 if event.is_none() {
274 self.bootstrap = None;
275 }
276 event
277 }
278
279 pub async fn recv(&mut self) -> Result<SessionEvent, RecvError> {
295 match self.next().await {
296 Some(Ok(event)) => Ok(event),
297 Some(Err(lagged)) => Err(lagged.into()),
298 None => Err(RecvErrorKind::Closed.into()),
299 }
300 }
301}
302
303impl Stream for EventSubscription {
304 type Item = Result<SessionEvent, Lagged>;
305
306 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
307 if let Some(event) = self.next_bootstrap_event() {
308 return Poll::Ready(Some(Ok(event)));
309 }
310 let Some(inner) = self.inner.as_mut() else {
311 return Poll::Ready(None);
312 };
313 match Pin::new(inner).poll_next(cx) {
314 Poll::Ready(Some(Ok(event))) => Poll::Ready(Some(Ok(event))),
315 Poll::Ready(Some(Err(BroadcastStreamRecvError::Lagged(n)))) => {
316 Poll::Ready(Some(Err(Lagged(n))))
317 }
318 Poll::Ready(None) => Poll::Ready(None),
319 Poll::Pending => Poll::Pending,
320 }
321 }
322}
323
324impl Drop for EventSubscription {
325 fn drop(&mut self) {
326 if let Some(bootstrap) = &self.bootstrap {
327 bootstrap.abandon();
328 }
329 }
330}
331
332#[must_use = "dropping the subscription unsubscribes"]
338pub struct LifecycleSubscription {
339 inner: BroadcastStream<SessionLifecycleEvent>,
340}
341
342impl LifecycleSubscription {
343 pub(crate) fn new(rx: Receiver<SessionLifecycleEvent>) -> Self {
344 Self {
345 inner: BroadcastStream::new(rx),
346 }
347 }
348
349 pub async fn recv(&mut self) -> Result<SessionLifecycleEvent, RecvError> {
364 match self.next().await {
365 Some(Ok(event)) => Ok(event),
366 Some(Err(lagged)) => Err(lagged.into()),
367 None => Err(RecvErrorKind::Closed.into()),
368 }
369 }
370}
371
372impl Stream for LifecycleSubscription {
373 type Item = Result<SessionLifecycleEvent, Lagged>;
374
375 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
376 match Pin::new(&mut self.inner).poll_next(cx) {
377 Poll::Ready(Some(Ok(event))) => Poll::Ready(Some(Ok(event))),
378 Poll::Ready(Some(Err(BroadcastStreamRecvError::Lagged(n)))) => {
379 Poll::Ready(Some(Err(Lagged(n))))
380 }
381 Poll::Ready(None) => Poll::Ready(None),
382 Poll::Pending => Poll::Pending,
383 }
384 }
385}
386
387#[cfg(test)]
388mod tests;