1use monoloop_connector::{
4 ConnectionEnd, ConnectionEndKind, Connector, OpenConnection, OpenedRawConnection,
5};
6use monoloop_contracts::{
7 CanonicalUnitEvent, ConnectionId, EffectiveConfig, EncodedExchange, ExchangeId,
8 ExchangeInputPolicy, InterpretationEnd, InterpretationEndKind, InterpretationId,
9 InterpretationLimits, OutboundDialectEncoder, TransactionId,
10};
11use monoloop_interpreter::{InterpreterFactory, StartInterpretation};
12use std::sync::Arc;
13use std::time::Duration;
14use tokio::runtime::Handle;
15use tokio::sync::{mpsc, oneshot};
16use tokio::task::JoinSet;
17
18use super::executor_spawn::try_spawn;
19
20pub struct ExchangeOutcome {
22 pub exchange_id: ExchangeId,
24 pub connection_id: ConnectionId,
26 pub interpretation_id: InterpretationId,
28 pub external_session_id: Option<monoloop_contracts::ExternalSessionId>,
30 pub units: Vec<CanonicalUnitEvent>,
32 pub connection_end: ConnectionEnd,
34 pub interpretation_end: InterpretationEnd,
36 pub failure: Option<ExchangeFailure>,
38}
39
40#[derive(Clone, Copy, Debug, PartialEq, Eq)]
42pub enum ExchangeFailure {
43 ChannelOpenFailed,
45 EncodingFailed,
47 ConnectorFailed,
49 InterpretationFailed,
51 Cancelled,
53 Terminated,
55 LimitExceeded,
57}
58
59pub struct ExchangeParams<'a> {
61 pub executor: &'a Handle,
63 pub transaction_id: TransactionId,
65 pub connector: &'a dyn Connector,
67 pub encoder: &'a dyn OutboundDialectEncoder,
69 pub interpreter: &'a dyn InterpreterFactory,
71 pub endpoint_ref: &'a str,
73 pub credential_ref: Option<&'a str>,
75 pub session_attachment: Option<Arc<monoloop_connector::SessionAttachment>>,
77 pub input: &'a monoloop_contracts::CanonicalInput,
79 pub config: &'a EffectiveConfig,
81 pub tools: &'a [monoloop_contracts::ToolSpec],
83 pub interpretation_limits: InterpretationLimits,
85 pub deadline: Duration,
87 pub cleanup_deadline: Duration,
89 pub max_encoded_exchange_bytes: usize,
91 pub unit_tx: Option<mpsc::Sender<CanonicalUnitEvent>>,
93 pub session_id_tx: Option<oneshot::Sender<monoloop_contracts::ExternalSessionId>>,
95}
96
97pub struct EncodedExchangeParams<'a> {
99 pub executor: &'a Handle,
101 pub transaction_id: TransactionId,
103 pub exchange_id: ExchangeId,
105 pub connector: &'a dyn Connector,
107 pub interpreter: &'a dyn InterpreterFactory,
109 pub endpoint_ref: &'a str,
111 pub credential_ref: Option<&'a str>,
113 pub session_attachment: Option<Arc<monoloop_connector::SessionAttachment>>,
115 pub encoded: EncodedExchange,
117 pub interpretation_limits: InterpretationLimits,
119 pub deadline: Duration,
121 pub cleanup_deadline: Duration,
123 pub max_encoded_exchange_bytes: usize,
125 pub unit_tx: Option<mpsc::Sender<CanonicalUnitEvent>>,
127}
128
129pub async fn run_exchange(params: ExchangeParams<'_>) -> Result<ExchangeOutcome, ExchangeFailure> {
131 let exchange_id = ExchangeId::generate();
132 let encoded = params
133 .encoder
134 .encode_initial(monoloop_contracts::InitialEncodeRequest {
135 transaction_id: ¶ms.transaction_id,
136 exchange_id: &exchange_id,
137 input: params.input,
138 config: params.config,
139 tools: params.tools,
140 })
141 .map_err(|_| ExchangeFailure::EncodingFailed)?;
142 if encoded.bytes.len() > params.max_encoded_exchange_bytes {
143 return Err(ExchangeFailure::EncodingFailed);
144 }
145
146 open_and_run(
147 params.executor,
148 exchange_id,
149 params.connector,
150 params.endpoint_ref,
151 params.credential_ref,
152 params.session_attachment,
153 encoded,
154 params.interpreter,
155 params.interpretation_limits,
156 params.deadline,
157 params.cleanup_deadline,
158 params.unit_tx,
159 params.session_id_tx,
160 )
161 .await
162}
163
164pub async fn run_encoded_exchange(
166 params: EncodedExchangeParams<'_>,
167) -> Result<ExchangeOutcome, ExchangeFailure> {
168 if params.encoded.bytes.len() > params.max_encoded_exchange_bytes {
169 return Err(ExchangeFailure::EncodingFailed);
170 }
171 open_and_run(
172 params.executor,
173 params.exchange_id,
174 params.connector,
175 params.endpoint_ref,
176 params.credential_ref,
177 params.session_attachment,
178 params.encoded,
179 params.interpreter,
180 params.interpretation_limits,
181 params.deadline,
182 params.cleanup_deadline,
183 params.unit_tx,
184 None,
185 )
186 .await
187}
188
189#[allow(clippy::too_many_arguments)]
190async fn open_and_run(
191 executor: &Handle,
192 exchange_id: ExchangeId,
193 connector: &dyn Connector,
194 endpoint_ref: &str,
195 credential_ref: Option<&str>,
196 session_attachment: Option<Arc<monoloop_connector::SessionAttachment>>,
197 encoded: EncodedExchange,
198 interpreter: &dyn InterpreterFactory,
199 interpretation_limits: InterpretationLimits,
200 deadline: Duration,
201 cleanup_deadline: Duration,
202 unit_tx: Option<mpsc::Sender<CanonicalUnitEvent>>,
203 session_id_tx: Option<oneshot::Sender<monoloop_contracts::ExternalSessionId>>,
204) -> Result<ExchangeOutcome, ExchangeFailure> {
205 let connection_id = ConnectionId::generate();
206 let interpretation_id = InterpretationId::generate();
207
208 let mut open = OpenConnection::new(connection_id.clone(), endpoint_ref);
209 open.credential_ref = credential_ref.map(|s| s.to_string());
210 if let Some(att) = session_attachment {
211 open = open.with_session_attachment(att);
212 }
213
214 let pending = connector.begin_open(open);
215 let mut open_guard = PendingOpenGuard {
217 control: Some(pending.control.clone()),
218 };
219 let opened = match tokio::time::timeout(deadline, pending.opened).await {
220 Ok(Ok(o)) => o,
221 Ok(Err(_)) => return Err(ExchangeFailure::ChannelOpenFailed),
222 Err(_) => return Err(ExchangeFailure::ChannelOpenFailed),
223 };
224 let _ = open_guard.control.take();
226
227 if let Some(tx) = session_id_tx {
228 if let Some(ref ext) = opened.external_session_id {
229 let _ = tx.send(ext.clone());
230 }
231 }
232
233 run_opened_exchange(
234 executor,
235 exchange_id,
236 interpretation_id,
237 opened,
238 encoded,
239 interpreter,
240 interpretation_limits,
241 deadline,
242 cleanup_deadline,
243 unit_tx,
244 )
245 .await
246}
247
248struct PendingOpenGuard {
250 control: Option<monoloop_connector::ConnectionControlHandle>,
251}
252
253impl Drop for PendingOpenGuard {
254 fn drop(&mut self) {
255 if let Some(ctrl) = self.control.take() {
256 let _ = ctrl.terminate(monoloop_connector::TerminationReason::CallerForced);
257 }
258 }
259}
260
261#[allow(clippy::too_many_arguments)]
262async fn run_opened_exchange(
263 executor: &Handle,
264 exchange_id: ExchangeId,
265 interpretation_id: InterpretationId,
266 opened: OpenedRawConnection,
267 encoded: EncodedExchange,
268 interpreter: &dyn InterpreterFactory,
269 limits: InterpretationLimits,
270 deadline: Duration,
271 cleanup_deadline: Duration,
272 unit_tx: Option<mpsc::Sender<CanonicalUnitEvent>>,
273) -> Result<ExchangeOutcome, ExchangeFailure> {
274 let join_grace = cleanup_deadline.max(Duration::from_millis(50));
275 let connection_id = opened.connection_id.clone();
276 let interpretation = interpreter
277 .start(StartInterpretation {
278 interpretation_id: interpretation_id.clone(),
279 connection_id: connection_id.clone(),
280 external_session_id: opened.external_session_id.clone(),
281 dialect: opened.dialect.clone(),
282 limits,
283 })
284 .map_err(|_| ExchangeFailure::InterpretationFailed)?;
285
286 let output = Arc::clone(&opened.output);
288 let interp_in = interpretation.input.clone();
289 let mut joins = JoinSet::new();
290 joins.spawn_on(
291 async move {
292 loop {
293 match output.receive().await {
294 Ok(Some(chunk)) => {
295 if interp_in.push_bytes(chunk).await.is_err() {
296 break;
297 }
298 }
299 Ok(None) => {
300 let _ = interp_in.finish_clean().await;
301 break;
302 }
303 Err(e) => {
304 use monoloop_contracts::ConnectorErrorKind;
305 match e.kind {
306 ConnectorErrorKind::Cancelled => {
307 let _ = interp_in.cancel().await;
308 }
309 ConnectorErrorKind::Terminated => {
310 let _ = interp_in.cancel().await;
311 }
312 _ => {
313 let _ = interp_in.transport_failed().await;
314 }
315 }
316 break;
317 }
318 }
319 }
320 },
321 executor,
322 );
323
324 if !encoded.bytes.is_empty() && opened.input.send(encoded.bytes.clone()).await.is_err() {
326 abort_joins(&mut joins).await;
327 return Err(ExchangeFailure::ConnectorFailed);
328 }
329 match encoded.input_policy {
330 ExchangeInputPolicy::SendAndFinish => {
331 if opened.input.finish().await.is_err() {
332 abort_joins(&mut joins).await;
333 return Err(ExchangeFailure::ConnectorFailed);
334 }
335 }
336 ExchangeInputPolicy::SendAndRetain => {}
337 }
338
339 let events_handle = interpretation.events;
343 let max_retained_units = 10_000usize;
344 let units = Arc::new(tokio::sync::Mutex::new(Vec::<CanonicalUnitEvent>::new()));
345 let retention_exceeded = Arc::new(std::sync::atomic::AtomicBool::new(false));
346 let units_task = {
347 let units = Arc::clone(&units);
348 let retention_exceeded = Arc::clone(&retention_exceeded);
349 let unit_tx = unit_tx;
350 try_spawn(executor, async move {
351 while let Some(ev) = events_handle.recv().await {
352 match ev {
353 monoloop_contracts::InterpreterOutputEvent::Unit(u) => {
354 let unit = *u;
355 if let Some(ref tx) = unit_tx {
356 if tx.send(unit.clone()).await.is_err() {
357 break;
358 }
359 }
360 let mut guard = units.lock().await;
361 if guard.len() >= max_retained_units {
362 retention_exceeded.store(true, std::sync::atomic::Ordering::SeqCst);
363 break;
364 }
365 guard.push(unit);
366 }
367 monoloop_contracts::InterpreterOutputEvent::Ended(_) => break,
368 }
369 }
370 })
371 .map_err(|_| ExchangeFailure::ConnectorFailed)?
372 };
373
374 let mut guard = ExchangeGuard {
377 control: Some(opened.control.clone()),
378 joins: Some(joins),
379 units_abort: Some(units_task.abort_handle()),
380 };
381
382 let completion = interpretation.completion;
383 let conn_completion = opened.completion;
384 let external_session_id = opened.external_session_id.clone();
385 let open_control = opened.control.clone();
386
387 let (interp_end, conn_end) = tokio::select! {
388 _ = tokio::time::sleep(deadline) => {
389 let _ = open_control
390 .terminate(monoloop_connector::TerminationReason::CallerForced);
391 if let Some(abort) = guard.units_abort.take() {
392 abort.abort();
393 }
394 let _ = tokio::time::timeout(join_grace, units_task).await;
395 if let Some(mut joins) = guard.joins.take() {
396 abort_joins(&mut joins).await;
397 }
398 let _ = guard.control.take();
399 return Err(ExchangeFailure::ConnectorFailed);
400 }
401 ends = async {
402 let i = completion.wait().await;
403 let c = conn_completion.wait().await;
404 (i, c)
405 } => ends,
406 };
407
408 if let Some(mut joins) = guard.joins.take() {
410 let _ = tokio::time::timeout(join_grace, async {
411 while joins.join_next().await.is_some() {}
412 })
413 .await;
414 abort_joins(&mut joins).await;
415 }
416 let mut units_task = units_task;
419 if let Some(abort) = guard.units_abort.take() {
420 match tokio::time::timeout(join_grace, &mut units_task).await {
421 Ok(_) => {}
422 Err(_) => {
423 abort.abort();
424 let _ = tokio::time::timeout(join_grace, units_task).await;
425 }
426 }
427 } else {
428 let _ = tokio::time::timeout(join_grace, units_task).await;
429 }
430 let _ = guard.control.take();
431
432 if retention_exceeded.load(std::sync::atomic::Ordering::SeqCst) {
433 return Err(ExchangeFailure::LimitExceeded);
434 }
435
436 let units = units.lock().await.clone();
437 let failure = reconcile_terminals(&conn_end, &interp_end);
438
439 Ok(ExchangeOutcome {
440 exchange_id,
441 connection_id,
442 interpretation_id,
443 external_session_id,
444 units,
445 connection_end: conn_end,
446 interpretation_end: interp_end,
447 failure,
448 })
449}
450
451struct ExchangeGuard {
453 control: Option<monoloop_connector::ConnectionControlHandle>,
454 joins: Option<JoinSet<()>>,
455 units_abort: Option<tokio::task::AbortHandle>,
456}
457
458impl Drop for ExchangeGuard {
459 fn drop(&mut self) {
460 if let Some(ctrl) = self.control.take() {
461 let _ = ctrl.terminate(monoloop_connector::TerminationReason::CallerForced);
462 }
463 if let Some(h) = self.units_abort.take() {
464 h.abort();
465 }
466 if let Some(mut joins) = self.joins.take() {
467 joins.abort_all();
468 }
469 }
470}
471
472fn reconcile_terminals(
473 conn: &ConnectionEnd,
474 interp: &InterpretationEnd,
475) -> Option<ExchangeFailure> {
476 match conn.kind {
477 ConnectionEndKind::Cancelled => return Some(ExchangeFailure::Cancelled),
478 ConnectionEndKind::Terminated => return Some(ExchangeFailure::Terminated),
479 ConnectionEndKind::TransportFailure => return Some(ExchangeFailure::ConnectorFailed),
480 ConnectionEndKind::RemoteEof | ConnectionEndKind::LocalShutdown => {}
481 }
482 match interp.kind {
483 InterpretationEndKind::Complete => None,
484 InterpretationEndKind::Cancelled => Some(ExchangeFailure::Cancelled),
485 InterpretationEndKind::Terminated => Some(ExchangeFailure::Terminated),
486 InterpretationEndKind::TransportFailed => Some(ExchangeFailure::ConnectorFailed),
487 InterpretationEndKind::DialectFailed
488 | InterpretationEndKind::LimitExceeded
489 | InterpretationEndKind::InvariantFailed => Some(ExchangeFailure::InterpretationFailed),
490 }
491}
492
493async fn abort_joins(joins: &mut JoinSet<()>) {
494 joins.abort_all();
495 while joins.join_next().await.is_some() {}
496}