feat: finalized support for chat

This commit is contained in:
Georg Bauer
2026-07-24 20:14:26 +02:00
parent 3154f57a2f
commit 6560f20129
9 changed files with 356 additions and 40 deletions

173
Cargo.lock generated
View File

@@ -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"

View File

@@ -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"

View File

@@ -0,0 +1 @@
ALTER TABLE sessions DROP COLUMN last_tokens_per_second;

View File

@@ -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);

View File

@@ -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<f32>,
#[cfg(target_os = "macos")]
generation_worker: Option<GenerationWorker>,
error: Option<String>,
@@ -418,6 +420,7 @@ pub(super) struct ChatMessage {
reasoning_complete: bool,
pub(super) reasoning_open: bool,
pub(super) content: String,
pub(super) markdown: Vec<markdown::Item>,
}
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<StoredMessage> 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<f32>,
},
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<GenerationWorker, String> {
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<Message> {
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());
}
}

View File

@@ -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]

View File

@@ -262,6 +262,7 @@ pub struct Session {
pub title: String,
pub context_used: i32,
pub context_limit: i32,
pub last_tokens_per_second: Option<f32>,
}
#[derive(Insertable)]
@@ -483,13 +484,18 @@ impl Database {
session_id: i32,
used: u32,
limit: u32,
tokens_per_second: Option<f32>,
) -> 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);

View File

@@ -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<f32>),
) -> 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<f32>),
) -> 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))
}

View File

@@ -63,6 +63,7 @@ diesel::table! {
title -> Text,
context_used -> Integer,
context_limit -> Integer,
last_tokens_per_second -> Nullable<Float>,
}
}