Skip to main content

chatty_rs/models/
conversation.rs

1#[cfg(test)]
2#[path = "conversation_test.rs"]
3mod tests;
4
5use crate::{
6    config,
7    config::constants::HELLO_MESSAGE,
8    models::{Message, message::Issuer},
9};
10
11#[derive(Debug, Clone)]
12pub struct Conversation {
13    id: String,
14    title: String,
15    messages: Vec<Message>,
16    contexts: Vec<Context>,
17    created_at: chrono::DateTime<chrono::Utc>,
18    updated_at: Option<chrono::DateTime<chrono::Utc>>,
19}
20
21impl Conversation {
22    pub fn new_hello() -> Self {
23        let mut conversation = Self::default();
24        conversation.messages.push(Message::new_system(
25            "system",
26            config::instance()
27                .general
28                .hello_message
29                .as_deref()
30                .unwrap_or(HELLO_MESSAGE),
31        ));
32        conversation
33    }
34
35    pub fn with_id(mut self, id: impl Into<String>) -> Self {
36        self.id = id.into();
37        self
38    }
39
40    pub fn with_created_at(mut self, timestamp: chrono::DateTime<chrono::Utc>) -> Self {
41        self.created_at = timestamp;
42        if self.updated_at.is_none() {
43            self.updated_at = Some(timestamp);
44        }
45        self
46    }
47
48    pub fn with_updated_at(mut self, timestamp: chrono::DateTime<chrono::Utc>) -> Self {
49        self.updated_at = Some(timestamp);
50        self
51    }
52
53    pub fn with_title(mut self, title: impl Into<String>) -> Self {
54        self.title = title.into();
55        self
56    }
57
58    pub fn set_updated_at(&mut self, timestamp: chrono::DateTime<chrono::Utc>) {
59        self.updated_at = Some(timestamp);
60    }
61
62    pub fn with_messages(mut self, messages: Vec<Message>) -> Self {
63        self.messages = messages;
64        self.messages.sort_by(|a, b| {
65            a.created_at()
66                .partial_cmp(&b.created_at())
67                .unwrap_or(std::cmp::Ordering::Equal)
68        });
69        self
70    }
71
72    pub fn with_context(mut self, context: Vec<Context>) -> Self {
73        self.contexts = context;
74        self.contexts.sort_by(|a, b| {
75            a.created_at()
76                .partial_cmp(&b.created_at())
77                .unwrap_or(std::cmp::Ordering::Equal)
78        });
79        self
80    }
81
82    pub fn set_id(&mut self, id: impl Into<String>) {
83        self.id = id.into();
84    }
85
86    pub fn set_title(&mut self, title: impl Into<String>) {
87        self.title = title.into();
88    }
89
90    pub fn append_message(&mut self, message: Message) {
91        self.messages.push(message);
92        self.messages.sort_by(|a, b| {
93            a.created_at()
94                .partial_cmp(&b.created_at())
95                .unwrap_or(std::cmp::Ordering::Equal)
96        });
97        self.updated_at = Some(chrono::Utc::now());
98    }
99
100    pub fn append_context(&mut self, context: Context) {
101        self.contexts.push(context);
102        self.contexts.sort_by(|a, b| {
103            a.created_at()
104                .partial_cmp(&b.created_at())
105                .unwrap_or(std::cmp::Ordering::Equal)
106        });
107    }
108
109    pub fn created_at(&self) -> chrono::DateTime<chrono::Utc> {
110        self.created_at
111    }
112
113    pub fn updated_at(&self) -> chrono::DateTime<chrono::Utc> {
114        self.updated_at.unwrap_or(self.created_at)
115    }
116
117    pub fn messages(&self) -> &[Message] {
118        &self.messages
119    }
120
121    pub fn title(&self) -> &str {
122        &self.title
123    }
124
125    pub fn id(&self) -> &str {
126        &self.id
127    }
128
129    pub fn last_message(&self) -> Option<&Message> {
130        self.messages.last()
131    }
132
133    pub fn last_mut_message(&mut self) -> Option<&mut Message> {
134        self.messages.last_mut()
135    }
136
137    pub fn len(&self) -> usize {
138        self.messages.len()
139    }
140
141    pub fn is_empty(&self) -> bool {
142        self.messages.is_empty()
143    }
144
145    pub fn messages_mut(&mut self) -> &mut Vec<Message> {
146        &mut self.messages
147    }
148
149    pub fn contexts_mut(&mut self) -> &mut Vec<Context> {
150        &mut self.contexts
151    }
152
153    pub fn contexts(&self) -> &[Context] {
154        &self.contexts
155    }
156
157    /// Return a vector of messages. The return vector is always end up
158    /// with a message from system
159    pub fn build_context(&self) -> Vec<Message> {
160        // If the conversation has less than 3 messages, return an empty vector
161        // 1 for hello message and 1 for user message so which means the conversation
162        // is not started yet. No context is needed.
163        if self.messages.len() < 3 && self.contexts.is_empty() {
164            return vec![];
165        }
166
167        let mut context: Vec<Message> = self.contexts.iter().map(Message::from).collect();
168
169        match self.contexts.last() {
170            Some(ctx) => {
171                // Find the index of the last message in the messages and
172                // append the next messages to the context
173                let last_message_index = self
174                    .messages
175                    .iter()
176                    .position(|msg| msg.id() == ctx.last_message_id())
177                    .unwrap_or(self.messages.len() - 2);
178                // Append the next messages to the context
179                context.extend(self.messages[last_message_index + 1..].to_vec());
180            }
181            None => context.extend(self.messages[1..].to_vec()),
182        }
183
184        if !context.last().unwrap().is_system() {
185            context.pop();
186        }
187
188        context
189    }
190
191    /// Calculate the total token count of the conversation.
192    /// This function will calculate the token count based on the context (if any)
193    /// and the messages started from the last context.
194    pub fn token_count(&self) -> usize {
195        let last_message_id = self
196            .contexts
197            .last()
198            .map(|ctx| ctx.last_message_id())
199            .unwrap_or_default();
200        if last_message_id.is_empty() {
201            return self.messages.iter().map(|msg| msg.token_count()).sum();
202        }
203
204        let tokens: usize = self.contexts.iter().map(|ctx| ctx.token_count()).sum();
205        let last_message_index = self
206            .messages
207            .iter()
208            .position(|msg| msg.id() == last_message_id)
209            .unwrap_or(self.messages.len() - 1);
210
211        let message_token = self
212            .messages
213            .iter()
214            .skip(last_message_index + 1)
215            .map(|msg| msg.token_count())
216            .sum::<usize>();
217        tokens + message_token
218    }
219}
220
221impl Default for Conversation {
222    fn default() -> Self {
223        Self {
224            id: "".to_string(),
225            title: "New Chat".to_string(),
226            messages: vec![],
227            contexts: vec![],
228            created_at: chrono::Utc::now(),
229            updated_at: None,
230        }
231    }
232}
233
234#[derive(Debug, Clone)]
235pub struct Context {
236    id: String,
237    content: String,
238    last_message_id: String,
239    token_count: usize,
240    created_at: chrono::DateTime<chrono::Utc>,
241}
242
243impl Context {
244    pub fn new(last_message_id: &str) -> Self {
245        Self {
246            id: uuid::Uuid::new_v4().to_string(),
247            content: String::new(),
248            token_count: 0,
249            last_message_id: last_message_id.to_string(),
250            created_at: chrono::Utc::now(),
251        }
252    }
253
254    pub fn with_token_count(mut self, token_count: usize) -> Self {
255        self.token_count = token_count;
256        self
257    }
258
259    pub fn with_id(mut self, id: impl Into<String>) -> Self {
260        self.id = id.into();
261        self
262    }
263
264    pub fn with_content(mut self, content: impl Into<String>) -> Self {
265        self.content = content.into();
266        self
267    }
268
269    pub fn with_created_at(mut self, timestamp: chrono::DateTime<chrono::Utc>) -> Self {
270        self.created_at = timestamp;
271        self
272    }
273
274    pub fn append_content(&mut self, content: impl Into<String>) {
275        self.content.push_str(&content.into());
276    }
277
278    pub fn id(&self) -> &str {
279        &self.id
280    }
281
282    pub fn content(&self) -> &str {
283        &self.content
284    }
285
286    pub fn last_message_id(&self) -> &str {
287        &self.last_message_id
288    }
289
290    pub fn created_at(&self) -> chrono::DateTime<chrono::Utc> {
291        self.created_at
292    }
293
294    pub fn token_count(&self) -> usize {
295        self.token_count
296    }
297
298    pub fn set_token_count(&mut self, token_count: usize) {
299        self.token_count = token_count;
300    }
301}
302
303impl From<&Context> for Message {
304    fn from(value: &Context) -> Message {
305        Message::new_system("system", &value.content)
306            .with_id(&value.id)
307            .with_created_at(value.created_at)
308            .with_token_count(value.token_count)
309            .with_context(true)
310    }
311}
312
313pub fn filter_issuer(issuer: Option<&Issuer>, msg: &Message) -> bool {
314    if issuer.is_none() {
315        return true;
316    }
317
318    let value;
319    let is_system = match issuer.unwrap() {
320        Issuer::System(sys) => {
321            value = sys.to_string();
322            true
323        }
324        Issuer::User(val) => {
325            value = val.to_string();
326            false
327        }
328    };
329
330    if is_system != msg.is_system() {
331        return false;
332    }
333
334    value.is_empty() || msg.issuer_str() == value
335}
336
337pub trait FindMessage {
338    fn last_message_of(&self, issuer: Option<Issuer>) -> Option<&Message>;
339    fn last_message_of_mut(&mut self, issuer: Option<Issuer>) -> Option<&mut Message>;
340}
341
342impl FindMessage for Vec<Message> {
343    fn last_message_of(&self, issuer: Option<Issuer>) -> Option<&Message> {
344        self.iter()
345            .rev()
346            .find(|&msg| filter_issuer(issuer.as_ref(), msg))
347    }
348
349    fn last_message_of_mut(&mut self, issuer: Option<Issuer>) -> Option<&mut Message> {
350        self.iter_mut()
351            .rev()
352            .find(|msg| filter_issuer(issuer.as_ref(), msg))
353    }
354}
355
356impl FindMessage for Conversation {
357    fn last_message_of(&self, issuer: Option<Issuer>) -> Option<&Message> {
358        self.messages.last_message_of(issuer)
359    }
360
361    fn last_message_of_mut(&mut self, issuer: Option<Issuer>) -> Option<&mut Message> {
362        self.messages.last_message_of_mut(issuer)
363    }
364}