1use std::{any::Any, cell::Cell, fmt, marker::PhantomData, rc::Rc};
2
3use futures::{FutureExt, future::LocalBoxFuture};
4
5use super::{
6 DiagnosticEvent, DiagnosticOutcome, DiagnosticSource, InvocationContext, NativeAppRuntime,
7 NativeStreamEndpointBinding, RequestPermit, RuntimeFailure, diagnostics::diagnostic_operation,
8};
9
10pub trait StreamCapability: 'static {
12 type OpenRequest: 'static;
14 type Message: 'static;
16 type DomainError: 'static;
18 const ID: &'static str;
20 const DESCRIPTOR_VERSION: &'static str;
22}
23
24#[derive(Clone, Debug, PartialEq)]
26pub enum StreamEvent<M, E> {
27 Message(M),
29 PeerHalfClosed,
31 Terminal(Result<(), E>),
33}
34
35pub enum NativeStreamItem {
37 Message(Box<dyn Any>),
39 PeerHalfClosed,
41 Terminal(Result<(), Box<dyn Any>>),
43}
44
45impl fmt::Debug for NativeStreamItem {
46 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
47 match self {
48 Self::Message(_) => formatter.write_str("Message(<erased>)"),
49 Self::PeerHalfClosed => formatter.write_str("PeerHalfClosed"),
50 Self::Terminal(Ok(())) => formatter.write_str("Terminal(Ok(()))"),
51 Self::Terminal(Err(_)) => formatter.write_str("Terminal(Err(<erased>))"),
52 }
53 }
54}
55
56pub trait NativeStreamSession: fmt::Debug {
58 fn send(&self, message: Box<dyn Any>) -> LocalBoxFuture<'static, Result<(), RuntimeFailure>>;
60 fn receive(&self) -> LocalBoxFuture<'static, Result<NativeStreamItem, RuntimeFailure>>;
62 fn close_send(&self) -> LocalBoxFuture<'static, Result<(), RuntimeFailure>>;
64 fn cancel(&self);
66}
67
68pub type NativeStreamOpenFuture = LocalBoxFuture<
70 'static,
71 Result<Result<Box<dyn NativeStreamSession>, Box<dyn Any>>, RuntimeFailure>,
72>;
73
74pub trait NativeStreamEndpoint: fmt::Debug {
76 fn capability_id(&self) -> &'static str;
78 fn descriptor_version(&self) -> &'static str;
80 fn operations(&self) -> &'static [&'static str];
82 fn open(
84 &self,
85 operation: &str,
86 request: Box<dyn Any>,
87 context: InvocationContext,
88 ) -> NativeStreamOpenFuture;
89}
90
91#[derive(Debug)]
93pub struct NativeStreamHandle<C: StreamCapability> {
94 endpoints: Vec<NativeStreamEndpointBinding>,
95 runtime: Rc<NativeAppRuntime>,
96 caller_instance: String,
97 allow_before_ready: bool,
98 capability: PhantomData<fn() -> C>,
99}
100
101impl<C: StreamCapability> Clone for NativeStreamHandle<C> {
102 fn clone(&self) -> Self {
103 Self {
104 endpoints: self.endpoints.clone(),
105 runtime: self.runtime.clone(),
106 caller_instance: self.caller_instance.clone(),
107 allow_before_ready: self.allow_before_ready,
108 capability: PhantomData,
109 }
110 }
111}
112
113impl<C: StreamCapability> NativeStreamHandle<C> {
114 pub(crate) fn from_endpoints(
115 endpoints: &[NativeStreamEndpointBinding],
116 runtime: Rc<NativeAppRuntime>,
117 caller_instance: &str,
118 allow_before_ready: bool,
119 ) -> Self {
120 Self {
121 endpoints: endpoints.to_vec(),
122 runtime,
123 caller_instance: caller_instance.to_owned(),
124 allow_before_ready,
125 capability: PhantomData,
126 }
127 }
128
129 pub fn binding_count(&self) -> usize {
131 self.endpoints.len()
132 }
133
134 pub async fn open(
136 &self,
137 operation: &str,
138 request: C::OpenRequest,
139 ) -> Result<Result<NativeStream<C>, C::DomainError>, RuntimeFailure> {
140 let context = self.next_context();
141 self.open_with_context(operation, context, request).await
142 }
143
144 pub async fn open_with_context(
146 &self,
147 operation: &str,
148 context: InvocationContext,
149 request: C::OpenRequest,
150 ) -> Result<Result<NativeStream<C>, C::DomainError>, RuntimeFailure> {
151 let context = context
152 .for_caller(&self.caller_instance)
153 .for_target(C::ID, operation);
154 let started_at = (self.runtime.driver.now)();
155 let operation_name = self
156 .endpoints
157 .first()
158 .and_then(|endpoint| diagnostic_operation(endpoint.state.operations, operation));
159 let request_id = context.request_id();
160 self.runtime
161 .diagnostics
162 .emit(DiagnosticSource::Invocation, started_at, |_| {
163 DiagnosticEvent::InvocationStarted {
164 requirement_id: self
165 .endpoints
166 .first()
167 .map(|endpoint| endpoint.requirement_id.clone()),
168 request_id,
169 caller_instance: Some(self.caller_instance.clone()),
170 provider_instance: self
171 .endpoints
172 .first()
173 .map(|endpoint| endpoint.plugin_instance.clone()),
174 capability: C::ID,
175 operation: operation_name,
176 }
177 });
178 let result = self
179 .open_with_context_inner(operation, context, request)
180 .await;
181 let outcome = match &result {
182 Ok(Ok(_)) => DiagnosticOutcome::Succeeded,
183 Ok(Err(_)) => DiagnosticOutcome::DomainError,
184 Err(error) => DiagnosticOutcome::RuntimeFailure(error.into()),
185 };
186 self.runtime.diagnostics.emit(
187 DiagnosticSource::Invocation,
188 (self.runtime.driver.now)(),
189 |_| DiagnosticEvent::InvocationCompleted {
190 requirement_id: self
191 .endpoints
192 .first()
193 .map(|endpoint| endpoint.requirement_id.clone()),
194 request_id,
195 caller_instance: Some(self.caller_instance.clone()),
196 provider_instance: self
197 .endpoints
198 .first()
199 .map(|endpoint| endpoint.plugin_instance.clone()),
200 capability: C::ID,
201 operation: operation_name,
202 outcome,
203 elapsed: (self.runtime.driver.now)().saturating_sub(started_at),
204 },
205 );
206 if let Err(error) = &result {
207 self.runtime.diagnostics.emit_runtime_failure(
208 (self.runtime.driver.now)(),
209 self.endpoints
210 .first()
211 .map(|endpoint| endpoint.plugin_instance.as_str()),
212 error,
213 );
214 }
215 result
216 }
217
218 async fn open_with_context_inner(
219 &self,
220 operation: &str,
221 context: InvocationContext,
222 request: C::OpenRequest,
223 ) -> Result<Result<NativeStream<C>, C::DomainError>, RuntimeFailure> {
224 if self.runtime.shutdown_started.get()
225 || (!self.allow_before_ready && self.runtime.admission.is_closed())
226 {
227 return Err(RuntimeFailure::AdmissionClosed);
228 }
229 let endpoint = match self.endpoints.as_slice() {
230 [] => return Err(RuntimeFailure::Unavailable { capability: C::ID }),
231 [endpoint] => endpoint,
232 endpoints => {
233 return Err(RuntimeFailure::AmbiguousBinding {
234 capability: C::ID,
235 providers: endpoints.len(),
236 });
237 }
238 };
239 let snapshot = endpoint
240 .state
241 .snapshot()
242 .ok_or(RuntimeFailure::Unavailable { capability: C::ID })?;
243 let admission = endpoint
244 .admission(operation)
245 .ok_or_else(|| RuntimeFailure::UnknownOperation {
246 capability: C::ID,
247 operation: operation.to_owned(),
248 })?
249 .clone();
250 let permit = admission
251 .acquire(C::ID, operation, &context, &self.runtime.driver)
252 .await?;
253 if !endpoint.state.is_current(snapshot.generation) {
254 return Err(RuntimeFailure::Unavailable { capability: C::ID });
255 }
256 let generation_cancellation = snapshot.cancellation.clone();
257 let endpoint_impl = snapshot.endpoint.clone();
258 let operation_name = operation.to_owned();
259 let (outcome, permit) = super::settlement::operation(
260 &self.runtime,
261 &endpoint.plugin_instance,
262 &context,
263 snapshot.cancellation,
264 C::ID,
265 move |execution_context| {
266 async move {
267 let outcome = endpoint_impl
268 .open(&operation_name, Box::new(request), execution_context)
269 .await;
270 outcome.map(|outcome| (outcome, permit))
271 }
272 .boxed_local()
273 },
274 )
275 .await??;
276 match outcome {
277 Ok(session) => Ok(Ok(NativeStream::new(
278 session,
279 self.runtime.clone(),
280 generation_cancellation,
281 endpoint.plugin_instance.clone(),
282 context,
283 permit,
284 ))),
285 Err(error) => Ok(Err(error
286 .downcast::<C::DomainError>()
287 .map(|error| *error)
288 .map_err(|_| RuntimeFailure::ProtocolViolation { capability: C::ID })?)),
289 }
290 }
291
292 fn next_context(&self) -> InvocationContext {
293 InvocationContext::new(
294 self.next_request_id(),
295 None,
296 super::CancellationToken::new(),
297 )
298 .with_caller_instance(self.caller_instance.clone())
299 }
300
301 fn next_request_id(&self) -> super::RequestId {
302 let request_id = self.runtime.request_ids.get();
303 self.runtime.request_ids.set(request_id.saturating_add(1));
304 request_id
305 }
306}
307
308#[derive(Debug)]
310pub struct NativeStream<C: StreamCapability> {
311 inner: Rc<dyn NativeStreamSession>,
312 runtime: Rc<NativeAppRuntime>,
313 generation_cancellation: super::CancellationToken,
314 plugin_instance: String,
315 context: InvocationContext,
316 _permit: RequestPermit,
317 local_half_closed: Cell<bool>,
318 peer_half_closed: Cell<bool>,
319 terminal_seen: Cell<bool>,
320 cancelled: Cell<bool>,
321 capability: PhantomData<fn() -> C>,
322}
323
324impl<C: StreamCapability> NativeStream<C> {
325 fn new(
326 session: Box<dyn NativeStreamSession>,
327 runtime: Rc<NativeAppRuntime>,
328 generation_cancellation: super::CancellationToken,
329 plugin_instance: String,
330 context: InvocationContext,
331 permit: RequestPermit,
332 ) -> Self {
333 Self {
334 inner: Rc::from(session),
335 runtime,
336 generation_cancellation,
337 plugin_instance,
338 context,
339 _permit: permit,
340 local_half_closed: Cell::new(false),
341 peer_half_closed: Cell::new(false),
342 terminal_seen: Cell::new(false),
343 cancelled: Cell::new(false),
344 capability: PhantomData,
345 }
346 }
347
348 pub async fn send(&self, message: C::Message) -> Result<(), RuntimeFailure> {
350 if let Some(error) = self.cancelled_outcome() {
351 return Err(error);
352 }
353 if self.local_half_closed.get() || self.terminal_seen.get() {
354 return Err(Self::protocol_violation());
355 }
356 let inner = self.inner.clone();
357 super::settlement::operation(
358 &self.runtime,
359 &self.plugin_instance,
360 &self.context,
361 self.generation_cancellation.clone(),
362 C::ID,
363 move |_| inner.send(Box::new(message)),
364 )
365 .await
366 .map_err(|error| self.finish_with_error(error))?
367 .map_err(|error| self.finish_with_error(error))
368 }
369
370 pub async fn receive(&self) -> Result<StreamEvent<C::Message, C::DomainError>, RuntimeFailure> {
372 if let Some(error) = self.cancelled_outcome() {
373 return Err(error);
374 }
375 if self.terminal_seen.get() {
376 return Err(Self::protocol_violation());
377 }
378 let inner = self.inner.clone();
379 let item = super::settlement::operation(
380 &self.runtime,
381 &self.plugin_instance,
382 &self.context,
383 self.generation_cancellation.clone(),
384 C::ID,
385 move |_| inner.receive(),
386 )
387 .await
388 .map_err(|error| self.finish_with_error(error))?
389 .map_err(|error| self.finish_with_error(error))?;
390 match item {
391 super::NativeStreamItem::Message(message) => {
392 if self.peer_half_closed.get() {
393 return Err(self.finish_with_error(Self::protocol_violation()));
394 }
395 message
396 .downcast::<C::Message>()
397 .map(|message| StreamEvent::Message(*message))
398 .map_err(|_| self.finish_with_error(Self::protocol_violation()))
399 }
400 super::NativeStreamItem::PeerHalfClosed => {
401 if self.peer_half_closed.replace(true) {
402 return Err(self.finish_with_error(Self::protocol_violation()));
403 }
404 Ok(StreamEvent::PeerHalfClosed)
405 }
406 super::NativeStreamItem::Terminal(outcome) => {
407 if self.terminal_seen.replace(true) {
408 return Err(self.finish_with_error(Self::protocol_violation()));
409 }
410 let outcome = match outcome {
411 Ok(()) => Ok(()),
412 Err(error) => Err(error
413 .downcast::<C::DomainError>()
414 .map(|error| *error)
415 .map_err(|_| self.finish_with_error(Self::protocol_violation()))?),
416 };
417 Ok(StreamEvent::Terminal(outcome))
418 }
419 }
420 }
421
422 pub async fn close_send(&self) -> Result<(), RuntimeFailure> {
424 if let Some(error) = self.cancelled_outcome() {
425 return Err(error);
426 }
427 if self.terminal_seen.get() || self.local_half_closed.replace(true) {
428 return Err(Self::protocol_violation());
429 }
430 let inner = self.inner.clone();
431 let result = super::settlement::operation(
432 &self.runtime,
433 &self.plugin_instance,
434 &self.context,
435 self.generation_cancellation.clone(),
436 C::ID,
437 move |_| inner.close_send(),
438 )
439 .await
440 .map_err(|error| self.finish_with_error(error))?
441 .map_err(|error| self.finish_with_error(error));
442 let resource_exhausted = result
443 .as_ref()
444 .err()
445 .is_some_and(|error| matches!(error, RuntimeFailure::ResourceExhausted { .. }));
446 if resource_exhausted {
447 self.local_half_closed.set(false);
448 }
449 result
450 }
451
452 pub fn cancel(&self) {
454 if !self.terminal_seen.get() && !self.cancelled.replace(true) {
455 self.context.cancellation().cancel();
456 self.inner.cancel();
457 }
458 }
459
460 pub const fn request_id(&self) -> super::RequestId {
462 self.context.request_id()
463 }
464
465 fn protocol_violation() -> RuntimeFailure {
466 RuntimeFailure::ProtocolViolation { capability: C::ID }
467 }
468
469 fn cancelled_outcome(&self) -> Option<RuntimeFailure> {
470 if !self.cancelled.get() {
471 return None;
472 }
473 if self.terminal_seen.replace(true) {
474 Some(Self::protocol_violation())
475 } else {
476 Some(RuntimeFailure::Cancelled {
477 request_id: self.context.request_id(),
478 })
479 }
480 }
481
482 fn finish_with_error(&self, error: RuntimeFailure) -> RuntimeFailure {
483 self.runtime.diagnostics.emit_runtime_failure(
484 (self.runtime.driver.now)(),
485 Some(&self.plugin_instance),
486 &error,
487 );
488 if !matches!(error, RuntimeFailure::ResourceExhausted { .. }) {
489 self.terminal_seen.set(true);
490 if !self.cancelled.replace(true) {
491 self.inner.cancel();
493 }
494 }
495 error
496 }
497}
498
499impl<C: StreamCapability> Drop for NativeStream<C> {
500 fn drop(&mut self) {
501 if !self.cancelled.replace(true) && !self.terminal_seen.get() {
502 self.inner.cancel();
503 }
504 }
505}
506
507pub type StreamSession<C> = NativeStream<C>;