Skip to main content

strop_engine/editor/picker/
ranking.rs

1//! Correlated picker ranking: native CPU actors publish through the same AppEvent
2//! and replay boundary as every other service. No scoring occurs on an input key.
3use super::{Editor, PickerId};
4use std::collections::HashSet;
5use std::sync::mpsc::{self, Receiver, Sender};
6use strop_core::worker::{FailureKind, Outcome, Ticket};
7use strop_picker::{RankingEvent, RankingWorker};
8
9#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
10pub struct Key {
11    pub picker: PickerId,
12    pub query: String,
13    pub items: usize,
14}
15#[derive(serde::Serialize, serde::Deserialize)]
16pub struct Event {
17    pub picker: PickerId,
18    pub update: RankingEvent<Key>,
19}
20pub(crate) struct State {
21    pub tx: Sender<Event>,
22    pub rx: Option<Receiver<Event>>,
23    pub retiring: HashSet<PickerId>,
24}
25impl Default for State {
26    fn default() -> Self {
27        let (tx, rx) = mpsc::channel();
28        Self {
29            tx,
30            rx: Some(rx),
31            retiring: HashSet::new(),
32        }
33    }
34}
35impl Editor {
36    pub(super) fn start_picker_ranking(&mut self) {
37        let Some(glue) = self.picker.as_mut() else {
38            return;
39        };
40        let id = glue.id;
41        glue.rank_alive = true;
42        match self.tape.request("picker.ranker", &id) {
43            Ok(false) => {
44                self.request_picker_ranking();
45                return;
46            }
47            Err(error) => {
48                glue.rank_alive = false;
49                glue.picker.error = Some(error.to_string());
50                return;
51            }
52            Ok(true) => {}
53        }
54        let tx = self.picker_ranking.tx.clone();
55        match RankingWorker::start(move |update| tx.send(Event { picker: id, update }).is_ok()) {
56            Ok(worker) => glue.rank_worker = Some(worker),
57            Err(error) => {
58                glue.rank_alive = false;
59                glue.picker.error = Some(format!("picker ranking unavailable: {error}"));
60                return;
61            }
62        }
63        self.request_picker_ranking();
64    }
65
66    pub(crate) fn request_picker_ranking(&mut self) {
67        let Some(glue) = self.picker.as_mut() else {
68            return;
69        };
70        if !glue.rank_alive {
71            return;
72        }
73        if glue.ranked_query.as_deref() != Some(glue.picker.input.text.as_str()) {
74            glue.picker.clear_results();
75        }
76        let request = match self.worker_ids.allocate() {
77            Ok(request) => request,
78            Err(error) => {
79                glue.picker.error = Some(error.message);
80                return;
81            }
82        };
83        let filter = glue.picker.filter_request();
84        let ticket = Ticket {
85            request,
86            key: Key {
87                picker: glue.id,
88                query: filter.query.clone(),
89                items: filter.catalog.len(),
90            },
91        };
92        glue.rank_pending = Some(ticket.clone());
93        match self.tape.request("picker.rank", &ticket) {
94            Ok(false) => return,
95            Ok(true) => {}
96            Err(error) => {
97                self.handle_picker_ranking(Event {
98                    picker: ticket.key.picker,
99                    update: RankingEvent::Completed(strop_core::worker::Completion {
100                        ticket,
101                        outcome: Outcome::failed(FailureKind::Protocol, error.to_string()),
102                    }),
103                });
104                return;
105            }
106        }
107        let result = glue
108            .rank_worker
109            .as_ref()
110            .ok_or_else(|| std::io::Error::other("picker ranking worker is absent"))
111            .and_then(|worker| worker.submit(ticket.clone(), filter));
112        if let Err(error) = result {
113            self.handle_picker_ranking(Event {
114                picker: ticket.key.picker,
115                update: RankingEvent::Completed(strop_core::worker::Completion {
116                    ticket,
117                    outcome: Outcome::failed(FailureKind::Disconnected, error.to_string()),
118                }),
119            });
120        }
121    }
122
123    pub(crate) fn handle_picker_ranking(&mut self, event: Event) {
124        if matches!(&event.update, RankingEvent::Stopped) {
125            self.picker_ranking.retiring.remove(&event.picker);
126        }
127        let Some(glue) = self.picker.as_mut().filter(|glue| glue.id == event.picker) else {
128            return;
129        };
130        match event.update {
131            RankingEvent::Completed(completion) => {
132                if glue.rank_pending.as_ref() != Some(&completion.ticket) {
133                    return;
134                }
135                glue.rank_pending = None;
136                match completion.outcome {
137                    Outcome::Success(ranking) => {
138                        if completion.ticket.key.query != glue.picker.input.text
139                            || !glue.picker.install_ranking(ranking)
140                        {
141                            return;
142                        }
143                        glue.ranked_query = Some(completion.ticket.key.query);
144                    }
145                    Outcome::Failed { failure, .. } => glue.picker.error = Some(failure.message),
146                    Outcome::Cancelled(_) => {}
147                }
148            }
149            RankingEvent::Stopped => {
150                glue.rank_alive = false;
151                if glue.rank_pending.take().is_some() {
152                    glue.picker.error = Some("picker ranking worker stopped".into());
153                    glue.revoke(strop_core::worker::CancelReason::OwnerClosed);
154                    glue.picker.streaming = false;
155                }
156            }
157        }
158        self.finish_pending_picker_accept();
159    }
160}