use aion_core::{ActivityId, RunId};
use aion_integrations::envelope_delta::EnvelopeDeltaDecoder;
use aion_proto::{
PerWorkflowSubscription, ProtoRunId, ProtoWorkflowId, StreamedActivityEvent,
TranscriptSubscription, WireError,
};
use aion_store::ActivityStreamKey;
use axum::extract::ws::{CloseFrame, Message, WebSocket, close_code};
use futures::{SinkExt, StreamExt};
use crate::activity_publisher::TranscriptStreamLagged;
use crate::error::ServerError;
use crate::namespace::{CallerIdentity, NamespaceOperation, SubscriptionScope, WorkflowTarget};
use crate::state::ServerState;
pub async fn serve_transcript_socket(
mut socket: WebSocket,
state: &ServerState,
caller: &CallerIdentity,
subscription: &TranscriptSubscription,
) -> Result<(), ServerError> {
let key = match authorize_transcript(state, caller, subscription).await {
Ok(key) => key,
Err(error) => {
super::socket::send_wire_error(&mut socket, &error.to_wire_error()).await?;
return Err(error);
}
};
let publisher = state.transcript_publisher();
let mut live = publisher.subscribe(key.clone(), subscription.after_seq);
let from_seq = subscription
.after_seq
.map_or(0, |seq| seq.saturating_add(1));
let replay = publisher
.replay_from(&key, from_seq)
.await
.map_err(ServerError::from)?;
let mut envelopes = EnvelopeDeltaDecoder::new();
for mut record in replay {
resolve_frame(&mut envelopes, &mut record.event);
if send_activity_frame(&mut socket, record.event)
.await?
.is_break()
{
return Ok(());
}
}
let (mut socket_tx, mut socket_rx) = socket.split();
loop {
tokio::select! {
client_message = socket_rx.next() => {
match client_message {
Some(Ok(Message::Close(_))) | None => {
return send_normal_close(&mut socket_tx).await;
}
Some(Ok(_other)) => {}
Some(Err(_error)) => return Ok(()),
}
}
item = live.next() => {
match item {
Some(Ok(mut event)) => {
resolve_frame(&mut envelopes, &mut event);
if forward_live_frame(&mut socket_tx, event).await?.is_break() {
return Ok(());
}
}
Some(Err(TranscriptStreamLagged { skipped })) => {
return deliver_transcript_terminal(&mut socket_tx, skipped).await;
}
None => return send_normal_close(&mut socket_tx).await,
}
}
}
}
}
fn resolve_frame(envelopes: &mut EnvelopeDeltaDecoder, event: &mut aion_core::ActivityEvent) {
if let Some(report) = crate::transcript_resolve::resolve_event(envelopes, event) {
crate::transcript_resolve::note_unresolved("ws:transcript", std::slice::from_ref(&report));
}
}
async fn authorize_transcript(
state: &ServerState,
caller: &CallerIdentity,
subscription: &TranscriptSubscription,
) -> Result<ActivityStreamKey, ServerError> {
let workflow_id = decode_workflow_id(subscription.workflow_id.as_ref())?;
let run_id = decode_run_id(subscription.run_id.as_ref())?;
let activity_id = decode_activity_id(subscription)?;
gate_transcript_workflow(state, caller, &subscription.namespace, &workflow_id).await?;
Ok(ActivityStreamKey::new(
workflow_id,
run_id,
activity_id,
subscription.attempt,
))
}
pub(crate) async fn gate_transcript_workflow(
state: &ServerState,
caller: &CallerIdentity,
namespace: &str,
workflow_id: &aion_core::WorkflowId,
) -> Result<(), ServerError> {
let per_workflow = PerWorkflowSubscription {
namespace: namespace.to_owned(),
workflow_id: Some(ProtoWorkflowId::from(workflow_id.clone())),
resume_from_seq: None,
};
let target = WorkflowTarget::workflow(workflow_id);
let scope = SubscriptionScope::PerWorkflow(&per_workflow, target);
let filter = aion::EventFilter {
workflow_id: Some(workflow_id.clone()),
..aion::EventFilter::default()
};
let operation = NamespaceOperation::subscribe(scope, &filter);
let scoped = state.namespace_guard().scope(caller, &operation).await?;
drop(scoped);
Ok(())
}
fn decode_workflow_id(
workflow_id: Option<&ProtoWorkflowId>,
) -> Result<aion_core::WorkflowId, ServerError> {
workflow_id
.cloned()
.ok_or_else(|| ServerError::Wire {
wire: WireError::invalid_input("transcript subscription workflow_id is missing"),
})?
.try_into()
.map_err(|wire| ServerError::Wire { wire })
}
fn decode_run_id(run_id: Option<&ProtoRunId>) -> Result<RunId, ServerError> {
run_id
.cloned()
.ok_or_else(|| ServerError::Wire {
wire: WireError::invalid_input("transcript subscription run_id is missing"),
})?
.try_into()
.map_err(|wire| ServerError::Wire { wire })
}
fn decode_activity_id(subscription: &TranscriptSubscription) -> Result<ActivityId, ServerError> {
let activity_id = subscription.activity_id.ok_or_else(|| ServerError::Wire {
wire: WireError::invalid_input("transcript subscription activity_id is missing"),
})?;
Ok(ActivityId::from(activity_id))
}
async fn send_activity_frame(
socket: &mut WebSocket,
event: aion_core::ActivityEvent,
) -> Result<std::ops::ControlFlow<()>, ServerError> {
let frame = encode_activity_frame(&event)?;
if socket.send(Message::Text(frame.into())).await.is_err() {
return Ok(std::ops::ControlFlow::Break(()));
}
Ok(std::ops::ControlFlow::Continue(()))
}
async fn forward_live_frame<Tx>(
socket_tx: &mut Tx,
event: aion_core::ActivityEvent,
) -> Result<std::ops::ControlFlow<()>, ServerError>
where
Tx: futures::Sink<Message> + Unpin,
<Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
{
let frame = match encode_activity_frame(&event) {
Ok(frame) => frame,
Err(error) => {
super::socket::send_wire_error(socket_tx, &error.to_wire_error()).await?;
return Err(error);
}
};
if socket_tx.send(Message::Text(frame.into())).await.is_err() {
return Ok(std::ops::ControlFlow::Break(()));
}
Ok(std::ops::ControlFlow::Continue(()))
}
fn encode_activity_frame(event: &aion_core::ActivityEvent) -> Result<String, ServerError> {
let frame = StreamedActivityEvent::new(event.clone());
serde_json::to_string(&frame).map_err(|source| ServerError::Wire {
wire: WireError::backend(format!(
"failed to serialize transcript event frame: {source}"
)),
})
}
async fn deliver_transcript_terminal<Tx>(
socket_tx: &mut Tx,
skipped: u64,
) -> Result<(), ServerError>
where
Tx: futures::Sink<Message> + Unpin,
<Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
{
let payload = serde_json::json!({
"error": { "code": "transcript_lagged", "skipped": skipped },
});
let payload = serde_json::to_string(&payload).map_err(|source| ServerError::Wire {
wire: WireError::backend(format!(
"failed to serialize transcript lag frame: {source}"
)),
})?;
if socket_tx.send(Message::Text(payload.into())).await.is_ok() {
let close = CloseFrame {
code: close_code::ERROR,
reason: "transcript_lagged".into(),
};
let close_result = socket_tx.send(Message::Close(Some(close))).await;
drop(close_result);
}
Err(ServerError::lagged_stream())
}
async fn send_normal_close<Tx>(socket_tx: &mut Tx) -> Result<(), ServerError>
where
Tx: futures::Sink<Message> + Unpin,
<Tx as futures::Sink<Message>>::Error: std::fmt::Debug,
{
let close = CloseFrame {
code: close_code::NORMAL,
reason: "subscription complete".into(),
};
let close_result = socket_tx.send(Message::Close(Some(close))).await;
drop(close_result);
Ok(())
}