Skip to main content

mj_controller/pollers/
types.rs

1use super::*;
2
3#[derive(Debug, Clone, Default)]
4pub struct QuotaRefreshBatch {
5    pub generation: u64,
6    /// The person asked for a refresh: probe every profile that is not on
7    /// hold, fresh or not. Without it, only profiles whose report is stale
8    /// (or missing) are probed.
9    pub refresh: bool,
10    pub profiles: Vec<QuotaRefreshRequest>,
11}
12
13#[derive(Debug)]
14pub enum QuotaUpdate {
15    Refreshing { profile_ids: Vec<String> },
16    Report(QuotaRefreshOutcome),
17    Finished { generation: u64 },
18}
19
20pub type WorkerPollTarget = RelaySessionTarget;
21pub type WorkerPollUpdate = SessionManagerUpdate;
22
23#[derive(Debug)]
24pub(super) struct WorkerDiagnosisEpisode {
25    pub(super) id: u64,
26    pub(super) error: String,
27    pub(super) diagnosed: bool,
28}
29
30#[derive(Debug, Default)]
31pub struct WorkerDiagnosisTracker {
32    pub(super) next_episode: u64,
33    pub(super) current: std::collections::BTreeMap<String, WorkerDiagnosisEpisode>,
34    pub(super) pending: std::collections::BTreeMap<String, u64>,
35}
36
37#[derive(Debug, Default, PartialEq, Eq)]
38pub struct WorkerDiagnosisCompletion {
39    pub display_error: Option<String>,
40    pub restart_episode: Option<u64>,
41}
42
43impl WorkerDiagnosisTracker {
44    pub fn observe(
45        &mut self,
46        session_id: &str,
47        connected: bool,
48        error: Option<String>,
49    ) -> Option<u64> {
50        if connected || error.is_none() {
51            self.current.remove(session_id);
52        }
53        let error = error?;
54        let episode = self
55            .current
56            .entry(session_id.to_owned())
57            .or_insert_with(|| {
58                self.next_episode = self.next_episode.wrapping_add(1).max(1);
59                WorkerDiagnosisEpisode {
60                    id: self.next_episode,
61                    error: error.clone(),
62                    diagnosed: false,
63                }
64            });
65        episode.error = error;
66        if episode.diagnosed || self.pending.contains_key(session_id) {
67            return None;
68        }
69        self.pending.insert(session_id.to_owned(), episode.id);
70        Some(episode.id)
71    }
72
73    pub fn finish(&mut self, session_id: &str, episode_id: u64) -> WorkerDiagnosisCompletion {
74        if self.pending.get(session_id) != Some(&episode_id) {
75            return WorkerDiagnosisCompletion::default();
76        }
77        self.pending.remove(session_id);
78        let Some(current) = self.current.get_mut(session_id) else {
79            return WorkerDiagnosisCompletion::default();
80        };
81        if current.id == episode_id {
82            current.diagnosed = true;
83            return WorkerDiagnosisCompletion {
84                display_error: Some(current.error.clone()),
85                restart_episode: None,
86            };
87        }
88        if !current.diagnosed {
89            self.pending.insert(session_id.to_owned(), current.id);
90            return WorkerDiagnosisCompletion {
91                display_error: None,
92                restart_episode: Some(current.id),
93            };
94        }
95        WorkerDiagnosisCompletion::default()
96    }
97}
98
99#[derive(Debug)]
100pub struct CapacityPollUpdate {
101    pub target_id: String,
102    pub result: std::result::Result<Option<DeploymentCapacityUsage>, String>,
103    pub sampled_at_epoch_seconds: u64,
104}
105
106pub fn projected_queued_prompts(
107    controller: &Controller,
108) -> Result<std::collections::BTreeMap<String, Vec<mj_core::relay::QueuedPrompt>>> {
109    let queues = crate::database::load_materialized_queued_prompts()?;
110    Ok(controller
111        .state
112        .sessions
113        .keys()
114        .filter_map(|session_id| {
115            queues
116                .get(session_id)
117                .map(|queue| (session_id.clone(), queued_prompt_entries(queue)))
118        })
119        .collect())
120}