From 6560f20129cfe729ea552d6000fb431cd62ebe2d Mon Sep 17 00:00:00 2001 From: Georg Bauer Date: Fri, 24 Jul 2026 20:14:26 +0200 Subject: [PATCH] feat: finalized support for chat --- Cargo.lock | 173 +++++++++++++++++- Cargo.toml | 2 +- .../down.sql | 1 + .../up.sql | 1 + src/app.rs | 123 +++++++++++-- src/app/view.rs | 64 +++++-- src/database.rs | 9 +- src/engine.rs | 22 ++- src/schema.rs | 1 + 9 files changed, 356 insertions(+), 40 deletions(-) create mode 100644 migrations/20260724234000_add_session_generation_speed/down.sql create mode 100644 migrations/20260724234000_add_session_generation_speed/up.sql diff --git a/Cargo.lock b/Cargo.lock index d744705..5f4740c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -340,6 +340,15 @@ version = "0.22.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" +[[package]] +name = "bincode" +version = "1.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b1f45e9417d87227c7a56d22e471c6206462cba514c7590c09aff4cf6d1ddcad" +dependencies = [ + "serde", +] + [[package]] name = "bit-set" version = "0.5.3" @@ -625,7 +634,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3538270d33cc669650c4b093848450d380def10c331d38c768e34cac80576e6e" dependencies = [ "termcolor", - "unicode-width", + "unicode-width 0.1.14", ] [[package]] @@ -1586,6 +1595,15 @@ dependencies = [ "windows-link", ] +[[package]] +name = "getopts" +version = "0.2.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfe4fbac503b8d1f88e6676011885f34b7174f46e59956bba534ba83abded4df" +dependencies = [ + "unicode-width 0.2.2", +] + [[package]] name = "getrandom" version = "0.2.17" @@ -1995,6 +2013,7 @@ checksum = "88acfabc84ec077eaf9ede3457ffa3a104626d79022a9bf7f296093b1d60c73f" dependencies = [ "iced_core", "iced_futures", + "iced_highlighter", "iced_renderer", "iced_widget", "iced_winit", @@ -2069,6 +2088,17 @@ dependencies = [ "unicode-segmentation", ] +[[package]] +name = "iced_highlighter" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bad88b25a1328cd4bb0b72d8e20f8207c0433649dc788f67e911423b9406f45c" +dependencies = [ + "iced_core", + "once_cell", + "syntect", +] + [[package]] name = "iced_renderer" version = "0.13.0" @@ -2139,13 +2169,16 @@ version = "0.13.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "81429e1b950b0e4bca65be4c4278fea6678ea782030a411778f26fa9f8983e1d" dependencies = [ + "iced_highlighter", "iced_renderer", "iced_runtime", "num-traits", "once_cell", + "pulldown-cmark", "rustc-hash 2.1.3", "thiserror 1.0.69", "unicode-segmentation", + "url", ] [[package]] @@ -2515,6 +2548,12 @@ dependencies = [ "x11", ] +[[package]] +name = "linked-hash-map" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0717cef1bc8b636c6e1c1bbdefc09e6322da8a9321966e8928ef80d20f7f770f" + [[package]] name = "linux-raw-sys" version = "0.4.15" @@ -3075,6 +3114,28 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +[[package]] +name = "onig" +version = "6.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cc3cbf698f9438986c11a880c90a6d04b9de27575afd28bbf45b154b6c709e2" +dependencies = [ + "bitflags 2.13.1", + "libc", + "once_cell", + "onig_sys", +] + +[[package]] +name = "onig_sys" +version = "69.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e68317604e77e53b85896388e1a803c1d21b74c899ec9e5e1112db90735edd7" +dependencies = [ + "cc", + "pkg-config", +] + [[package]] name = "orbclient" version = "0.3.55" @@ -3332,6 +3393,19 @@ version = "0.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b4596b6d070b27117e987119b4dac604f3c58cfb0b191112e24771b2faeac1a6" +[[package]] +name = "plist" +version = "1.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7da1d65da6dd5d1e44199ac0f58712d241c0f439f80adea8924d832384087f85" +dependencies = [ + "base64", + "indexmap", + "quick-xml 0.41.0", + "serde", + "time", +] + [[package]] name = "png" version = "0.17.16" @@ -3463,6 +3537,25 @@ version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3d595e54a326bc53c1c197b32d295e14b169e3cfeaa8dc82b529f947fba6bcf5" +[[package]] +name = "pulldown-cmark" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "679341d22c78c6c649893cbd6c3278dcbe9fc4faa62fea3a9296ae2b50c14625" +dependencies = [ + "bitflags 2.13.1", + "getopts", + "memchr", + "pulldown-cmark-escape", + "unicase", +] + +[[package]] +name = "pulldown-cmark-escape" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "007d8adb5ddab6f8e3f491ac63566a7d5002cc7ed73901f72057943fa71ae1ae" + [[package]] name = "quick-xml" version = "0.39.4" @@ -3472,6 +3565,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "quick-xml" +version = "0.41.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e660451e55124f798a69a5af3f49ccfbefbd41910eefd25caf2393e1f3473ec1" +dependencies = [ + "memchr", +] + [[package]] name = "quote" version = "1.0.47" @@ -3647,6 +3749,12 @@ dependencies = [ "thiserror 1.0.69", ] +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + [[package]] name = "renderdoc-sys" version = "1.1.0" @@ -3923,6 +4031,19 @@ dependencies = [ "syn 3.0.3", ] +[[package]] +name = "serde_json" +version = "1.0.151" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c841b55ecdae098c80dcae9cf767f6f8a0c2cdb3416bbef72181df4d0fe73f14" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + [[package]] name = "serde_repr" version = "0.1.21" @@ -4287,6 +4408,27 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "syntect" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "656b45c05d95a5704399aeef6bd0ddec7b2b3531b7c9e900abbf7c4d2190c925" +dependencies = [ + "bincode", + "flate2", + "fnv", + "once_cell", + "onig", + "plist", + "regex-syntax", + "serde", + "serde_derive", + "serde_json", + "thiserror 2.0.19", + "walkdir", + "yaml-rust", +] + [[package]] name = "sys-locale" version = "0.3.2" @@ -4644,6 +4786,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "unicase" +version = "2.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142" + [[package]] name = "unicode-bidi" version = "0.3.18" @@ -4704,6 +4852,12 @@ version = "0.1.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7dd6e30e90baa6f72411720665d41d89b9a3d039dc45b8faea1ddd07f617f6af" +[[package]] +name = "unicode-width" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4ac048d71ede7ee76d585517add45da530660ef4390e49b098733c6e897f254" + [[package]] name = "unicode-xid" version = "0.2.6" @@ -5045,7 +5199,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9c324a910fd86ebdc364a3e61ec1f11737d3b1d6c273c0239ee8ff4bc0d24b4a" dependencies = [ "proc-macro2", - "quick-xml", + "quick-xml 0.39.4", "quote", ] @@ -5556,6 +5710,15 @@ version = "0.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ec7a2a501ed189703dba8b08142f057e887dfc4b2cc4db2d343ac6376ba3e0b9" +[[package]] +name = "yaml-rust" +version = "0.4.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "56c1936c4cc7a1c9ab21a1ebb602eb942ba868cbd44a99cb7cdc5892335e1c85" +dependencies = [ + "linked-hash-map", +] + [[package]] name = "yazi" version = "0.1.6" @@ -5794,6 +5957,12 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "zmij" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" + [[package]] name = "zvariant" version = "4.2.0" diff --git a/Cargo.toml b/Cargo.toml index 6b54a98..6fb84ba 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -14,7 +14,7 @@ cc = "1.3.0" [dependencies] diesel = { version = "2.3.11", features = ["sqlite", "returning_clauses_for_sqlite_3_35", "64-column-tables"] } diesel_migrations = "2.3.2" -iced = { version = "0.13.1", features = ["svg", "tokio"] } +iced = { version = "0.13.1", features = ["highlighter", "markdown", "svg", "tokio"] } memmap2 = "0.9.11" png = "0.17.16" rfd = "0.15.4" diff --git a/migrations/20260724234000_add_session_generation_speed/down.sql b/migrations/20260724234000_add_session_generation_speed/down.sql new file mode 100644 index 0000000..dca1d53 --- /dev/null +++ b/migrations/20260724234000_add_session_generation_speed/down.sql @@ -0,0 +1 @@ +ALTER TABLE sessions DROP COLUMN last_tokens_per_second; diff --git a/migrations/20260724234000_add_session_generation_speed/up.sql b/migrations/20260724234000_add_session_generation_speed/up.sql new file mode 100644 index 0000000..1500ef2 --- /dev/null +++ b/migrations/20260724234000_add_session_generation_speed/up.sql @@ -0,0 +1 @@ +ALTER TABLE sessions ADD COLUMN last_tokens_per_second FLOAT CHECK (last_tokens_per_second IS NULL OR last_tokens_per_second >= 0); diff --git a/src/app.rs b/src/app.rs index 5d4d859..0da67d4 100644 --- a/src/app.rs +++ b/src/app.rs @@ -11,6 +11,7 @@ use crate::settings::{ RuntimePreferences, SpeculativePreferences, SsdPreferences, SteeringPreferences, StreamingCacheBudget, }; +use iced::widget::{markdown, scrollable}; use iced::{Size, Subscription, Task, keyboard, window}; use rfd::AsyncFileDialog; use std::fs; @@ -405,6 +406,7 @@ pub(crate) struct App { pub(super) generating: bool, pub(super) context_used: u32, pub(super) context_limit: u32, + pub(super) tokens_per_second: Option, #[cfg(target_os = "macos")] generation_worker: Option, error: Option, @@ -418,6 +420,7 @@ pub(super) struct ChatMessage { reasoning_complete: bool, pub(super) reasoning_open: bool, pub(super) content: String, + pub(super) markdown: Vec, } impl ChatMessage { @@ -429,18 +432,32 @@ impl ChatMessage { self.content.push_str(chunk); } } + + fn refresh_markdown(&mut self) { + if !self.user { + let content = if self.reasoning.is_some() { + self.content.trim_start() + } else { + &self.content + }; + self.markdown = markdown::parse(content).collect(); + } + } } impl From for ChatMessage { fn from(message: StoredMessage) -> Self { - Self { + let mut message = Self { id: message.id, user: message.user, reasoning: message.reasoning, reasoning_complete: message.reasoning_complete, reasoning_open: false, content: message.content, - } + markdown: Vec::new(), + }; + message.refresh_markdown(); + message } } @@ -466,8 +483,15 @@ enum GenerationCommand { #[cfg(target_os = "macos")] enum GenerationEvent { Loading, - Chunk { reasoning: bool, content: String }, - Context { used: u32, limit: u32 }, + Chunk { + reasoning: bool, + content: String, + }, + Context { + used: u32, + limit: u32, + tokens_per_second: Option, + }, Finished(Result<(), String>), } @@ -551,6 +575,7 @@ pub(crate) enum Message { DownloadProgressTick, ComposerChanged(String), ToggleReasoning(usize), + OpenLink(markdown::Url), SubmitPrompt, StopGeneration, GenerationTick, @@ -600,6 +625,7 @@ impl App { generating: false, context_used: 0, context_limit, + tokens_per_second: None, #[cfg(target_os = "macos")] generation_worker: None, error: None, @@ -639,6 +665,7 @@ impl App { generating: false, context_used: 0, context_limit, + tokens_per_second: None, #[cfg(target_os = "macos")] generation_worker: None, error: Some(format!("Could not open the project database: {error}")), @@ -915,7 +942,17 @@ impl App { message.reasoning_open = !message.reasoning_open; } } - Message::SubmitPrompt => self.start_generation(), + Message::OpenLink(url) => { + if matches!(url.scheme(), "http" | "https") + && let Err(error) = std::process::Command::new("open").arg(url.as_str()).spawn() + { + self.error = Some(format!("Could not open the link: {error}")); + } + } + Message::SubmitPrompt => { + self.start_generation(); + return scroll_chat_to_end(); + } Message::StopGeneration => { #[cfg(target_os = "macos")] if let Some(cancel) = self @@ -926,7 +963,11 @@ impl App { cancel.store(true, Ordering::Relaxed); } } - Message::GenerationTick => self.poll_generation(), + Message::GenerationTick => { + if self.poll_generation() { + return scroll_chat_to_end(); + } + } Message::ChooseProjectFolder => { self.choosing_folder = true; return Task::perform( @@ -965,6 +1006,7 @@ impl App { self.composer.clear(); self.context_used = 0; self.context_limit = self.preferences.context_tokens.max(0) as u32; + self.tokens_per_second = None; self.error = None; } Message::DeleteProject(project_id) => { @@ -996,6 +1038,7 @@ impl App { self.conversation.clear(); self.composer.clear(); self.context_used = 0; + self.tokens_per_second = None; } self.reload_projects(); } @@ -1018,7 +1061,13 @@ impl App { .iter() .flat_map(|project| &project.sessions) .find(|session| session.id == session_id) - .map(|session| (session.context_used, session.context_limit)); + .map(|session| { + ( + session.context_used, + session.context_limit, + session.last_tokens_per_second, + ) + }); let Some(database) = &mut self.database else { return Task::none(); }; @@ -1028,14 +1077,16 @@ impl App { self.composer.clear(); self.selected_project = Some(project_id); self.selected_session = Some(session_id); - let (used, limit) = saved_context.unwrap_or_default(); + let (used, limit, tokens_per_second) = saved_context.unwrap_or_default(); self.context_used = used.max(0) as u32; self.context_limit = if limit > 0 { limit as u32 } else { self.preferences.context_tokens.max(0) as u32 }; + self.tokens_per_second = tokens_per_second; self.error = None; + return scroll_chat_to_end(); } Err(error) => { self.error = Some(format!("Could not load the chat session: {error}")); @@ -1057,6 +1108,7 @@ impl App { self.conversation.clear(); self.composer.clear(); self.context_used = 0; + self.tokens_per_second = None; } self.reload_projects(); } @@ -1321,6 +1373,7 @@ impl App { self.composer.clear(); self.context_used = 0; self.context_limit = self.preferences.context_tokens.max(0) as u32; + self.tokens_per_second = None; self.error = None; self.reload_projects(); } @@ -1499,6 +1552,7 @@ impl App { self.conversation.push(user); self.conversation.push(assistant); self.generating = true; + self.tokens_per_second = None; self.error = None; } #[cfg(not(target_os = "macos"))] @@ -1509,11 +1563,11 @@ impl App { } } - fn poll_generation(&mut self) { + fn poll_generation(&mut self) -> bool { #[cfg(target_os = "macos")] let Some(worker) = &mut self.generation_worker else { self.generating = false; - return; + return false; }; #[cfg(target_os = "macos")] let mut transcript_changed = false; @@ -1531,9 +1585,14 @@ impl App { transcript_changed = true; } } - Ok(GenerationEvent::Context { used, limit }) => { + Ok(GenerationEvent::Context { + used, + limit, + tokens_per_second, + }) => { self.context_used = used; self.context_limit = limit; + self.tokens_per_second = tokens_per_second; context_changed = true; } Ok(GenerationEvent::Finished(result)) => { @@ -1554,6 +1613,10 @@ impl App { } } #[cfg(target_os = "macos")] + if transcript_changed && let Some(message) = self.conversation.last_mut() { + message.refresh_markdown(); + } + #[cfg(target_os = "macos")] if transcript_changed && let Some(message) = self.conversation.last() && let Some(database) = &mut self.database @@ -1578,9 +1641,12 @@ impl App { && let Some(session_id) = self.selected_session && let Some(database) = &mut self.database { - if let Err(error) = - database.update_session_context(session_id, self.context_used, self.context_limit) - { + if let Err(error) = database.update_session_context( + session_id, + self.context_used, + self.context_limit, + self.tokens_per_second, + ) { self.error = Some(format!("Could not save context usage: {error}")); } else if let Some(session) = self .projects @@ -1590,8 +1656,13 @@ impl App { { session.context_used = self.context_used as i32; session.context_limit = self.context_limit as i32; + session.last_tokens_per_second = self.tokens_per_second; } } + #[cfg(target_os = "macos")] + return transcript_changed; + #[cfg(not(target_os = "macos"))] + false } } @@ -1640,9 +1711,12 @@ fn spawn_generation_worker() -> Result { let _ = event_sender .send(GenerationEvent::Chunk { reasoning, content }); }, - |used, limit| { - let _ = - event_sender.send(GenerationEvent::Context { used, limit }); + |used, limit, tokens_per_second| { + let _ = event_sender.send(GenerationEvent::Context { + used, + limit, + tokens_per_second, + }); }, ); let _ = event_sender.send(GenerationEvent::Finished(result)); @@ -1706,6 +1780,14 @@ fn models_path() -> PathBuf { application_support_path().join("models") } +pub(super) fn chat_scroll_id() -> scrollable::Id { + scrollable::Id::new("chat-transcript") +} + +fn scroll_chat_to_end() -> Task { + scrollable::snap_to(chat_scroll_id(), scrollable::RelativeOffset::END) +} + fn session_checkpoint_path(session_id: i32) -> PathBuf { application_support_path() .join("kv-cache") @@ -1774,11 +1856,14 @@ mod tests { reasoning_complete: false, reasoning_open: true, content: String::new(), + markdown: Vec::new(), }; message.append(true, "working it out"); - message.append(false, "final answer"); + message.append(false, "**final answer**"); + message.refresh_markdown(); assert_eq!(message.reasoning.as_deref(), Some("working it out")); assert!(message.reasoning_complete); - assert_eq!(message.content, "final answer"); + assert_eq!(message.content, "**final answer**"); + assert!(!message.markdown.is_empty()); } } diff --git a/src/app/view.rs b/src/app/view.rs index bc28891..ec5074e 100644 --- a/src/app/view.rs +++ b/src/app/view.rs @@ -1,4 +1,6 @@ -use super::{ActiveDownload, App, Message, ModelDownload, ModelOperation, models_path}; +use super::{ + ActiveDownload, App, Message, ModelDownload, ModelOperation, chat_scroll_id, models_path, +}; use crate::database::{ProjectWithSessions, Session}; use crate::model::{ self, DownloadPhase, MODEL_CHOICES, ManagedArtifact, ManagedArtifactState, ModelChoice, @@ -6,8 +8,8 @@ use crate::model::{ use crate::settings::{GIB, REASONING_MODES}; use iced::theme::{Palette, palette}; use iced::widget::{ - Button, Space, Svg, button, checkbox, column, container, horizontal_rule, opaque, pick_list, - progress_bar, row, scrollable, stack, svg, text, text_input, + Button, Space, Svg, button, checkbox, column, container, horizontal_rule, markdown, opaque, + pick_list, progress_bar, row, scrollable, stack, svg, text, text_input, }; use iced::{Alignment, Background, Border, Color, Element, Length, Theme, window}; use std::path::Path; @@ -234,6 +236,7 @@ impl App { .spacing(8), ); } else { + let markdown_style = markdown::Style::from_palette(app_theme().palette()); for (index, message) in self.conversation.iter().enumerate() { let label = if message.user { "You" } else { "DS4" }; let active = self.generating && index + 1 == self.conversation.len(); @@ -267,20 +270,32 @@ impl App { } } if !message.content.is_empty() { - let content = if message.reasoning.is_some() { - message.content.trim_start() + if message.user || message.markdown.is_empty() { + let content = if message.reasoning.is_some() { + message.content.trim_start() + } else { + &message.content + }; + body = body.push(text(content).size(14)); } else { - &message.content - }; - body = body.push(text(content).size(14)); + body = body.push( + markdown::view( + &message.markdown, + markdown::Settings::with_text_size(14), + markdown_style, + ) + .map(Message::OpenLink), + ); + } } else if active && message.reasoning.is_none() { body = body.push(text("Loading model…").size(14)); } + let user = message.user; messages = messages.push( container(body) .padding(14) .width(Length::Fill) - .style(overview_style), + .style(move |theme| chat_message_style(theme, user)), ); } } @@ -304,7 +319,9 @@ impl App { self.context_used.min(self.context_limit) as f32 / self.context_limit as f32 }; let conversation = column![ - scrollable(messages).height(Length::Fill), + scrollable(messages) + .id(chat_scroll_id()) + .height(Length::Fill), container( column![ composer, @@ -312,10 +329,14 @@ impl App { row![ icon(ICON_PAPERCLIP, 19), text(format!( - "{} / {} tokens ({:.0}%)", + "{} / {} tokens ({:.0}%) • {}", self.context_used, self.context_limit, - context_fraction * 100.0 + context_fraction * 100.0, + self.tokens_per_second.map_or_else( + || "— tok/s".to_owned(), + |speed| format!("{speed:.1} tok/s") + ) )) .size(11) .color(muted_text()), @@ -1180,6 +1201,21 @@ fn overview_style(_: &Theme) -> container::Style { } } +fn chat_message_style(theme: &Theme, user: bool) -> container::Style { + if !user { + return overview_style(theme); + } + container::Style { + background: Some(Background::Color(Color::from_rgb8(27, 34, 44))), + border: Border { + color: Color::from_rgb8(54, 68, 88), + width: 1.0, + radius: 14.0.into(), + }, + ..container::Style::default() + } +} + fn muted_text() -> Color { Color::from_rgb8(174, 174, 178) } @@ -1207,6 +1243,10 @@ mod tests { let palette = theme.extended_palette(); assert_eq!(palette.background.weak.color, Color::from_rgb8(38, 38, 40)); assert_eq!(palette.secondary.base.color, Color::from_rgb8(47, 47, 50)); + assert_ne!( + chat_message_style(&theme, true).background, + chat_message_style(&theme, false).background + ); } #[test] diff --git a/src/database.rs b/src/database.rs index 46d0a80..12458c7 100644 --- a/src/database.rs +++ b/src/database.rs @@ -262,6 +262,7 @@ pub struct Session { pub title: String, pub context_used: i32, pub context_limit: i32, + pub last_tokens_per_second: Option, } #[derive(Insertable)] @@ -483,13 +484,18 @@ impl Database { session_id: i32, used: u32, limit: u32, + tokens_per_second: Option, ) -> Result<(), String> { let used = i32::try_from(used).map_err(|_| "Used context is too large to save")?; let limit = i32::try_from(limit).map_err(|_| "Context limit is too large to save")?; + if tokens_per_second.is_some_and(|speed| !speed.is_finite() || speed < 0.0) { + return Err("Generation speed is invalid".into()); + } diesel::update(sessions::table.find(session_id)) .set(( sessions::context_used.eq(used), sessions::context_limit.eq(limit), + sessions::last_tokens_per_second.eq(tokens_per_second), )) .execute(&mut self.connection) .map(|_| ()) @@ -669,7 +675,7 @@ mod tests { .update_message(assistant.id, Some("Reasoning"), true, "Answer") .unwrap(); database - .update_session_context(session.id, 1_234, 65_536) + .update_session_context(session.id, 1_234, 65_536, Some(12.5)) .unwrap(); drop(database); @@ -677,6 +683,7 @@ mod tests { let projects = reopened.load_projects().unwrap(); assert_eq!(projects[0].sessions[0].context_used, 1_234); assert_eq!(projects[0].sessions[0].context_limit, 65_536); + assert_eq!(projects[0].sessions[0].last_tokens_per_second, Some(12.5)); let messages = reopened.load_messages(session.id).unwrap(); assert_eq!(messages.len(), 2); assert!(messages[0].user); diff --git a/src/engine.rs b/src/engine.rs index 3fe6e76..fd59fd2 100644 --- a/src/engine.rs +++ b/src/engine.rs @@ -10,6 +10,8 @@ use gguf::{F16, F32, Gguf, I32, IQ2_XXS, Q2_K, Q4_0, Q4_K, Q5_K, Q6_K, Q8_0, Ten use sha2::{Digest, Sha256}; use std::path::{Path, PathBuf}; use std::sync::atomic::{AtomicBool, Ordering}; +#[cfg(target_os = "macos")] +use std::time::Instant; use tokenizer::Tokenizer; #[cfg(target_os = "macos")] @@ -334,7 +336,7 @@ impl Generator { settings: &TurnSettings, cancelled: &AtomicBool, mut emit: impl FnMut(bool, String), - mut progress: impl FnMut(u32, u32), + mut progress: impl FnMut(u32, u32, Option), ) -> Result<(), String> { self.select_checkpoint(checkpoint)?; let (generated, prompt_complete) = @@ -370,7 +372,7 @@ impl Generator { settings: &TurnSettings, cancelled: &AtomicBool, emit: &mut impl FnMut(bool, String), - progress: &mut impl FnMut(u32, u32), + progress: &mut impl FnMut(u32, u32, Option), ) -> Result<(ChatTurn, bool), String> { let tokens = match messages.split_last() { Some((latest, history)) @@ -407,7 +409,7 @@ impl Generator { )); } let reused = self.executor.align_prompt(&tokens)?; - progress(self.executor.position(), self.executor.context()); + progress(self.executor.position(), self.executor.context(), None); let mut rng = Rng::new(settings.seed.unwrap_or(0x4453_3453_4552_5645)); let mut reasoning = settings.reasoning_mode != ReasoningMode::Direct; let mut generated = ChatTurn { @@ -423,9 +425,11 @@ impl Generator { } self.executor.eval(token)?; if (index + 1).is_multiple_of(16) || index + 1 == prompt_tokens { - progress(self.executor.position(), self.executor.context()); + progress(self.executor.position(), self.executor.context(), None); } } + let generation_started = Instant::now(); + let mut generated_tokens = 0_u32; for _ in 0..settings .max_generated_tokens .max(0) @@ -469,7 +473,15 @@ impl Generator { emit(reasoning, content); } self.executor.eval(token)?; - progress(self.executor.position(), self.executor.context()); + generated_tokens += 1; + progress( + self.executor.position(), + self.executor.context(), + Some( + generated_tokens as f32 + / generation_started.elapsed().as_secs_f32().max(1.0e-6), + ), + ); } Ok((generated, true)) } diff --git a/src/schema.rs b/src/schema.rs index 6c901e3..ec3b8cf 100644 --- a/src/schema.rs +++ b/src/schema.rs @@ -63,6 +63,7 @@ diesel::table! { title -> Text, context_used -> Integer, context_limit -> Integer, + last_tokens_per_second -> Nullable, } }