Skip to main content

trustee_api/
state.rs

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