Skip to main content

graphflow_stream/
stream.rs

1#![doc = include_str!("../README.md")]
2
3use graph_flow::{
4    Context, ExecutionResult, ExecutionStatus, FlowRunner, Graph, InMemorySessionStorage,
5    NextAction, Session, SessionStorage, Task, TaskResult,
6    error::{GraphError, Result},
7};
8use std::future::Future;
9use std::sync::Arc;
10use std::sync::atomic::{AtomicU64, Ordering};
11use tokio::sync::mpsc;
12use tokio::task::JoinHandle;
13
14pub type StreamSender = mpsc::Sender<StreamEvent>;
15pub type StreamReceiver = mpsc::Receiver<StreamEvent>;
16
17#[derive(Debug, Clone, PartialEq, Eq)]
18pub enum StreamEvent {
19    TaskStarted { task_id: String },
20    Token { task_id: String, delta: String },
21    TaskFinished { task_id: String },
22    TaskFailed { task_id: String, error: String },
23}
24
25tokio::task_local! {
26    static STREAM_TX: StreamSender;
27}
28
29pub async fn emit(event: StreamEvent) {
30    if let Ok(tx) = STREAM_TX.try_with(|tx| tx.clone()) {
31        let _ = tx.send(event).await;
32    }
33}
34
35pub async fn emit_started(task_id: impl Into<String>) {
36    emit(StreamEvent::TaskStarted {
37        task_id: task_id.into(),
38    })
39    .await;
40}
41
42pub async fn emit_token(task_id: impl Into<String>, delta: impl Into<String>) {
43    emit(StreamEvent::Token {
44        task_id: task_id.into(),
45        delta: delta.into(),
46    })
47    .await;
48}
49
50pub async fn emit_finished(task_id: impl Into<String>) {
51    emit(StreamEvent::TaskFinished {
52        task_id: task_id.into(),
53    })
54    .await;
55}
56
57pub async fn emit_failed(task_id: impl Into<String>, error: impl Into<String>) {
58    emit(StreamEvent::TaskFailed {
59        task_id: task_id.into(),
60        error: error.into(),
61    })
62    .await;
63}
64
65pub async fn forward_text_stream<S>(task_id: impl Into<String>, mut stream: S)
66where
67    S: futures_core::Stream<Item = String> + Unpin,
68{
69    use tokio_stream::StreamExt;
70
71    let task_id = task_id.into();
72    while let Some(delta) = stream.next().await {
73        emit_token(task_id.clone(), delta).await;
74    }
75    emit_finished(task_id).await;
76}
77
78pub fn run_streaming<F>(buffer: usize, fut: F) -> (StreamReceiver, JoinHandle<F::Output>)
79where
80    F: Future + Send + 'static,
81    F::Output: Send + 'static,
82{
83    let (tx, rx) = mpsc::channel(buffer);
84    let handle = tokio::spawn(STREAM_TX.scope(tx, fut));
85    (rx, handle)
86}
87
88pub fn spawn_task<T>(
89    task: Arc<T>,
90    context: Context,
91    buffer: usize,
92) -> (StreamReceiver, JoinHandle<Result<TaskResult>>)
93where
94    T: Task + Send + Sync + 'static,
95{
96    run_streaming(buffer, async move { task.run(context).await })
97}
98
99pub fn spawn_graph(
100    flow_runner: FlowRunner,
101    session_id: impl Into<String>,
102    buffer: usize,
103) -> (StreamReceiver, JoinHandle<Result<ExecutionResult>>) {
104    let session_id = session_id.into();
105    run_streaming(buffer, async move {
106        run_to_completion(&flow_runner, &session_id).await
107    })
108}
109
110async fn run_to_completion(flow_runner: &FlowRunner, session_id: &str) -> Result<ExecutionResult> {
111    loop {
112        let result = flow_runner.run(session_id).await?;
113        if !matches!(result.status, ExecutionStatus::Paused { .. }) {
114            return Ok(result);
115        }
116    }
117}
118
119static SUBGRAPH_SESSION_COUNTER: AtomicU64 = AtomicU64::new(0);
120
121pub struct SubgraphTask {
122    id: String,
123    graph: Arc<Graph>,
124    storage: Arc<dyn SessionStorage>,
125}
126
127impl SubgraphTask {
128    pub fn new(id: impl Into<String>, graph: Arc<Graph>) -> Self {
129        Self {
130            id: id.into(),
131            graph,
132            storage: Arc::new(InMemorySessionStorage::new()),
133        }
134    }
135
136    pub fn with_storage(mut self, storage: Arc<dyn SessionStorage>) -> Self {
137        self.storage = storage;
138        self
139    }
140}
141
142#[async_trait::async_trait]
143impl Task for SubgraphTask {
144    fn id(&self) -> &str {
145        &self.id
146    }
147
148    async fn run(&self, context: Context) -> Result<TaskResult> {
149        let start_task_id = self.graph.start_task_id().ok_or_else(|| {
150            GraphError::TaskNotFound(format!("subgraph '{}' has no start task", self.graph.id))
151        })?;
152
153        let suffix = SUBGRAPH_SESSION_COUNTER.fetch_add(1, Ordering::Relaxed);
154        let session_id = format!("{}:{}", self.id, suffix);
155
156        let mut session = Session::new_from_task(session_id.clone(), start_task_id)
157            .with_graph_id(self.graph.id.clone());
158        session.context = context;
159        self.storage.save(session).await?;
160
161        let runner = FlowRunner::new(self.graph.clone(), self.storage.clone());
162        let result = run_to_completion(&runner, &session_id).await?;
163
164        let next_action = match result.status {
165            ExecutionStatus::WaitingForInput => NextAction::WaitForInput,
166            _ => NextAction::Continue,
167        };
168        Ok(TaskResult::new(result.response, next_action))
169    }
170}