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