feat: finalized support for chat
This commit is contained in:
123
src/app.rs
123
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<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());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user