Skip to main content

trustee_api/
state.rs

1//! Shared server state: wraps `Session` and a broadcast channel.
2
3use std::sync::Arc;
4
5use tokio::sync::{broadcast, mpsc, Mutex};
6use trustee_core::session::Session;
7use trustee_core::types::TuiMessage;
8
9/// Shared state accessible by all axum handlers.
10#[derive(Clone)]
11pub struct ServerState {
12    /// The agent session, protected by a mutex.
13    pub session: Arc<Mutex<Session>>,
14    /// Broadcast sender for WebSocket fan-out.
15    /// Messages are JSON-serialized `TuiMessage` strings.
16    pub ws_tx: broadcast::Sender<String>,
17}
18
19impl ServerState {
20    /// Create new shared state from a session and broadcast sender.
21    pub fn new(session: Session, ws_tx: broadcast::Sender<String>) -> Self {
22        Self {
23            session: Arc::new(Mutex::new(session)),
24            ws_tx,
25        }
26    }
27
28    /// Spawn a background task that owns the workflow receiver and broadcasts
29    /// each message to all WebSocket subscribers.
30    ///
31    /// The receiver is moved into the task — no locking needed to await it.
32    /// When a message arrives, the task briefly locks the session to call
33    /// `handle_workflow_message`, then broadcasts the JSON to WebSocket clients.
34    pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
35        tokio::spawn(async move {
36            while let Some(msg) = workflow_rx.recv().await {
37                // Process the message through Session's handler (updates state)
38                {
39                    let mut session = self.session.lock().await;
40                    session.handle_workflow_message(msg.clone());
41                }
42
43                // Broadcast the raw message to WebSocket clients
44                let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
45                let _ = self.ws_tx.send(json);
46            }
47        });
48    }
49}
50
51/// Wrapper to serialize `TuiMessage` as JSON with a `type` discriminator.
52struct SerializableMessage<'a>(&'a TuiMessage);
53
54impl<'a> serde::Serialize for SerializableMessage<'a> {
55    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
56    where
57        S: serde::Serializer,
58    {
59        use serde::ser::SerializeStruct;
60
61        match self.0 {
62            TuiMessage::OutputLine(line) => {
63                let mut s = serializer.serialize_struct("msg", 2)?;
64                s.serialize_field("type", "OutputLine")?;
65                s.serialize_field("line", line)?;
66                s.end()
67            }
68            TuiMessage::StreamDelta(delta) => {
69                let mut s = serializer.serialize_struct("msg", 2)?;
70                s.serialize_field("type", "StreamDelta")?;
71                s.serialize_field("delta", delta)?;
72                s.end()
73            }
74            TuiMessage::ReasoningDelta(delta) => {
75                let mut s = serializer.serialize_struct("msg", 2)?;
76                s.serialize_field("type", "ReasoningDelta")?;
77                s.serialize_field("delta", delta)?;
78                s.end()
79            }
80            TuiMessage::WorkflowCompleted => {
81                let mut s = serializer.serialize_struct("msg", 2)?;
82                s.serialize_field("type", "WorkflowCompleted")?;
83                s.serialize_field("state", "Idle")?;
84                s.end()
85            }
86            TuiMessage::WorkflowError(err) => {
87                let mut s = serializer.serialize_struct("msg", 2)?;
88                s.serialize_field("type", "WorkflowError")?;
89                s.serialize_field("error", err)?;
90                s.end()
91            }
92            TuiMessage::ResumeInfo(_) => {
93                let mut s = serializer.serialize_struct("msg", 2)?;
94                s.serialize_field("type", "ResumeInfo")?;
95                s.serialize_field("state", "Idle")?;
96                s.end()
97            }
98            TuiMessage::TodoUpdate(content) => {
99                let mut s = serializer.serialize_struct("msg", 2)?;
100                s.serialize_field("type", "TodoUpdate")?;
101                s.serialize_field("content", content)?;
102                s.end()
103            }
104            TuiMessage::WorkflowCancelled => {
105                let mut s = serializer.serialize_struct("msg", 2)?;
106                s.serialize_field("type", "WorkflowCancelled")?;
107                s.serialize_field("state", "Idle")?;
108                s.end()
109            }
110            TuiMessage::HandoffReady(_) => {
111                let mut s = serializer.serialize_struct("msg", 2)?;
112                s.serialize_field("type", "HandoffReady")?;
113                s.serialize_field("state", "Idle")?;
114                s.end()
115            }
116            TuiMessage::ToolPending { tool_name, hint } => {
117                let mut s = serializer.serialize_struct("msg", 3)?;
118                s.serialize_field("type", "ToolPending")?;
119                s.serialize_field("tool_name", tool_name)?;
120                s.serialize_field("hint", hint)?;
121                s.end()
122            }
123            TuiMessage::ToolDone { tool_name, success, hint } => {
124                let mut s = serializer.serialize_struct("msg", 4)?;
125                s.serialize_field("type", "ToolDone")?;
126                s.serialize_field("tool_name", tool_name)?;
127                s.serialize_field("success", success)?;
128                s.serialize_field("hint", hint)?;
129                s.end()
130            }
131            TuiMessage::ContextTokensUpdated(count) => {
132                let mut s = serializer.serialize_struct("msg", 2)?;
133                s.serialize_field("type", "ContextTokensUpdated")?;
134                s.serialize_field("count", count)?;
135                s.end()
136            }
137            TuiMessage::McpServerStatus { name, connected, tool_count, error } => {
138                let mut s = serializer.serialize_struct("msg", 5)?;
139                s.serialize_field("type", "McpServerStatus")?;
140                s.serialize_field("name", name)?;
141                s.serialize_field("connected", connected)?;
142                s.serialize_field("tool_count", tool_count)?;
143                s.serialize_field("error", error)?;
144                s.end()
145            }
146        }
147    }
148}